mathcraft-ocr 0.2.7__tar.gz → 0.2.8__tar.gz
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.
Potentially problematic release.
This version of mathcraft-ocr might be problematic. Click here for more details.
- {mathcraft_ocr-0.2.7/mathcraft_ocr.egg-info → mathcraft_ocr-0.2.8}/PKG-INFO +2 -2
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/README_MATHCRAFT_OCR.md +1 -1
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/__init__.py +1 -1
- mathcraft_ocr-0.2.8/mathcraft_ocr/adapters/common.py +108 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/adapters/text_recognizer.py +19 -14
- mathcraft_ocr-0.2.8/mathcraft_ocr/devices.py +360 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/hardware.py +28 -27
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/providers.py +43 -4
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/runtime.py +8 -3
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/serialization.py +4 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8/mathcraft_ocr.egg-info}/PKG-INFO +2 -2
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr.egg-info/SOURCES.txt +1 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/pyproject.toml +1 -1
- mathcraft_ocr-0.2.7/mathcraft_ocr/adapters/common.py +0 -54
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/LICENSE +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/MANIFEST.in +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/__main__.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/adapters/__init__.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/adapters/formula_detector.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/adapters/formula_recognizer.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/adapters/text_detector.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/api.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/cache.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/cli.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/debug_blocks.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/doctor.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/downloader.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/error_patterns.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/errors.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/formula_lines.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/image.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/latex_alignment.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/latex_quality.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/layout.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/manifest.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/manifests/models.v1.json +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/profiles.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/results.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr/worker.py +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr.egg-info/dependency_links.txt +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr.egg-info/entry_points.txt +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr.egg-info/requires.txt +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/mathcraft_ocr.egg-info/top_level.txt +0 -0
- {mathcraft_ocr-0.2.7 → mathcraft_ocr-0.2.8}/setup.cfg +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: mathcraft-ocr
|
|
3
|
-
Version: 0.2.
|
|
3
|
+
Version: 0.2.8
|
|
4
4
|
Summary: ONNX-only OCR runtime for mathematical documents
|
|
5
5
|
Author: SakuraMathcraft
|
|
6
6
|
License-Expression: GPL-3.0-only
|
|
@@ -194,7 +194,7 @@ Model artifacts are downloaded from the MathCraft-Models release assets declared
|
|
|
194
194
|
- `cpu`: force CPU.
|
|
195
195
|
- `gpu`: request CUDA-capable ONNX Runtime.
|
|
196
196
|
|
|
197
|
-
The actual provider is available on results through the `provider` field.
|
|
197
|
+
The actual provider is available on recognition results through the `provider` field. Doctor and warmup reports also expose `device_id`, `device_name`, `device_uuid`, and `device_verified` under `provider_info`. GPU sessions bind the reported `device_id` explicitly; `device_verified` becomes true only after the runtime confirms the device used by initialized inference sessions.
|
|
198
198
|
|
|
199
199
|
## Development
|
|
200
200
|
|
|
@@ -148,7 +148,7 @@ Model artifacts are downloaded from the MathCraft-Models release assets declared
|
|
|
148
148
|
- `cpu`: force CPU.
|
|
149
149
|
- `gpu`: request CUDA-capable ONNX Runtime.
|
|
150
150
|
|
|
151
|
-
The actual provider is available on results through the `provider` field.
|
|
151
|
+
The actual provider is available on recognition results through the `provider` field. Doctor and warmup reports also expose `device_id`, `device_name`, `device_uuid`, and `device_verified` under `provider_info`. GPU sessions bind the reported `device_id` explicitly; `device_verified` becomes true only after the runtime confirms the device used by initialized inference sessions.
|
|
152
152
|
|
|
153
153
|
## Development
|
|
154
154
|
|
|
@@ -0,0 +1,108 @@
|
|
|
1
|
+
# coding: utf-8
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import importlib
|
|
6
|
+
from functools import lru_cache
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
|
|
9
|
+
from ..providers import GPU_PROVIDER_NAMES, ProviderInfo
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
ProviderOptions = tuple[tuple[str, str], ...]
|
|
13
|
+
ProviderConfig = tuple[str, ProviderOptions]
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def _ort():
|
|
17
|
+
return importlib.import_module("onnxruntime")
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def session_providers(provider_info: ProviderInfo) -> list[str | tuple[str, dict[str, int]]]:
|
|
21
|
+
available = list(provider_info.available_providers)
|
|
22
|
+
active = provider_info.active_provider
|
|
23
|
+
if active and active in GPU_PROVIDER_NAMES and "CPUExecutionProvider" in available:
|
|
24
|
+
return [(active, {"device_id": int(provider_info.device_id or 0)}), "CPUExecutionProvider"]
|
|
25
|
+
if active and active in GPU_PROVIDER_NAMES:
|
|
26
|
+
return [(active, {"device_id": int(provider_info.device_id or 0)})]
|
|
27
|
+
if "CPUExecutionProvider" in available:
|
|
28
|
+
return ["CPUExecutionProvider"]
|
|
29
|
+
return available
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def create_session(model_path: str | Path, provider_info: ProviderInfo):
|
|
33
|
+
model_path = str(Path(model_path).resolve())
|
|
34
|
+
provider_config = _freeze_provider_config(session_providers(provider_info))
|
|
35
|
+
session = _create_session_cached(model_path, provider_config)
|
|
36
|
+
enforce_session_provider(session, provider_info)
|
|
37
|
+
return session
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
@lru_cache(maxsize=16)
|
|
41
|
+
def _create_session_cached(model_path: str, providers: tuple[ProviderConfig, ...]):
|
|
42
|
+
ort = _ort()
|
|
43
|
+
configured = [
|
|
44
|
+
(name, {key: value for key, value in options}) if options else name
|
|
45
|
+
for name, options in providers
|
|
46
|
+
]
|
|
47
|
+
return ort.InferenceSession(
|
|
48
|
+
model_path,
|
|
49
|
+
providers=configured,
|
|
50
|
+
enable_fallback=False,
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _freeze_provider_config(
|
|
55
|
+
providers: list[str | tuple[str, dict[str, int]]],
|
|
56
|
+
) -> tuple[ProviderConfig, ...]:
|
|
57
|
+
frozen: list[ProviderConfig] = []
|
|
58
|
+
for item in providers:
|
|
59
|
+
if isinstance(item, str):
|
|
60
|
+
frozen.append((item, ()))
|
|
61
|
+
continue
|
|
62
|
+
name, options = item
|
|
63
|
+
frozen.append((name, tuple(sorted((str(key), str(value)) for key, value in options.items()))))
|
|
64
|
+
return tuple(frozen)
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def enforce_session_provider(session, provider_info: ProviderInfo) -> None:
|
|
68
|
+
validate_session_provider(
|
|
69
|
+
session,
|
|
70
|
+
str(provider_info.active_provider or ""),
|
|
71
|
+
int(provider_info.device_id or 0),
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def validate_session_provider(
|
|
76
|
+
session,
|
|
77
|
+
active_provider: str,
|
|
78
|
+
device_id: int = 0,
|
|
79
|
+
*,
|
|
80
|
+
runtime_name: str = "ONNX Runtime",
|
|
81
|
+
) -> None:
|
|
82
|
+
actual = list(session.get_providers() or [])
|
|
83
|
+
active = str(active_provider or "")
|
|
84
|
+
if not actual or actual[0] != active:
|
|
85
|
+
provider_kind = "ONNX GPU provider" if active in GPU_PROVIDER_NAMES else "ONNX provider"
|
|
86
|
+
raise RuntimeError(
|
|
87
|
+
f"requested {provider_kind} {active or '<none>'}, "
|
|
88
|
+
f"but {runtime_name} session providers are {actual}"
|
|
89
|
+
)
|
|
90
|
+
if active in GPU_PROVIDER_NAMES:
|
|
91
|
+
get_options = getattr(session, "get_provider_options", None)
|
|
92
|
+
if callable(get_options):
|
|
93
|
+
options_by_provider = get_options() or {}
|
|
94
|
+
active_options = options_by_provider.get(active, {}) if isinstance(options_by_provider, dict) else {}
|
|
95
|
+
actual_device_id = active_options.get("device_id") if isinstance(active_options, dict) else None
|
|
96
|
+
if actual_device_id not in (None, "") and int(actual_device_id) != int(device_id):
|
|
97
|
+
raise RuntimeError(
|
|
98
|
+
f"requested ONNX provider {active} device_id={int(device_id)}, "
|
|
99
|
+
f"but {runtime_name} session uses device_id={actual_device_id}"
|
|
100
|
+
)
|
|
101
|
+
disable_fallback = getattr(session, "disable_fallback", None)
|
|
102
|
+
if not callable(disable_fallback):
|
|
103
|
+
raise RuntimeError(f"{runtime_name} session cannot disable provider fallback")
|
|
104
|
+
disable_fallback()
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def clear_session_cache() -> None:
|
|
108
|
+
_create_session_cached.cache_clear()
|
|
@@ -10,6 +10,8 @@ from rapidocr import EngineType, LangRec, ModelType, OCRVersion
|
|
|
10
10
|
from rapidocr.ch_ppocr_rec import TextRecInput, TextRecognizer
|
|
11
11
|
from rapidocr.utils.typings import TaskType
|
|
12
12
|
|
|
13
|
+
from .common import validate_session_provider
|
|
14
|
+
|
|
13
15
|
|
|
14
16
|
class _Config(dict):
|
|
15
17
|
def __init__(self, *args, **kwargs):
|
|
@@ -55,11 +57,13 @@ def recognize_pp_text_lines(
|
|
|
55
57
|
def _create_pp_text_recognizer(model_dir: Path, provider_info) -> TextRecognizer:
|
|
56
58
|
model_dir = model_dir.resolve()
|
|
57
59
|
active_provider = str(getattr(provider_info, "active_provider", "") or "")
|
|
60
|
+
device_id = int(getattr(provider_info, "device_id", 0) or 0)
|
|
58
61
|
use_cuda = active_provider in {"CUDAExecutionProvider", "TensorrtExecutionProvider"}
|
|
59
62
|
use_dml = active_provider == "DmlExecutionProvider"
|
|
60
63
|
return _create_pp_text_recognizer_cached(
|
|
61
64
|
str(model_dir),
|
|
62
65
|
active_provider,
|
|
66
|
+
device_id,
|
|
63
67
|
use_cuda,
|
|
64
68
|
use_dml,
|
|
65
69
|
)
|
|
@@ -69,6 +73,7 @@ def _create_pp_text_recognizer(model_dir: Path, provider_info) -> TextRecognizer
|
|
|
69
73
|
def _create_pp_text_recognizer_cached(
|
|
70
74
|
model_dir: str,
|
|
71
75
|
active_provider: str,
|
|
76
|
+
device_id: int,
|
|
72
77
|
use_cuda: bool,
|
|
73
78
|
use_dml: bool,
|
|
74
79
|
) -> TextRecognizer:
|
|
@@ -103,13 +108,13 @@ def _create_pp_text_recognizer_cached(
|
|
|
103
108
|
"cpu_ep_cfg": {"arena_extend_strategy": "kSameAsRequested"},
|
|
104
109
|
"use_cuda": use_cuda,
|
|
105
110
|
"cuda_ep_cfg": {
|
|
106
|
-
"device_id":
|
|
111
|
+
"device_id": device_id,
|
|
107
112
|
"arena_extend_strategy": "kNextPowerOfTwo",
|
|
108
113
|
"cudnn_conv_algo_search": "EXHAUSTIVE",
|
|
109
114
|
"do_copy_in_default_stream": True,
|
|
110
115
|
},
|
|
111
116
|
"use_dml": use_dml,
|
|
112
|
-
"dm_ep_cfg":
|
|
117
|
+
"dm_ep_cfg": {"device_id": device_id},
|
|
113
118
|
"use_cann": False,
|
|
114
119
|
"cann_ep_cfg": {
|
|
115
120
|
"device_id": 0,
|
|
@@ -122,25 +127,25 @@ def _create_pp_text_recognizer_cached(
|
|
|
122
127
|
},
|
|
123
128
|
})
|
|
124
129
|
recognizer = TextRecognizer(config)
|
|
125
|
-
_enforce_strict_provider(recognizer, active_provider)
|
|
130
|
+
_enforce_strict_provider(recognizer, active_provider, device_id)
|
|
126
131
|
return recognizer
|
|
127
132
|
|
|
128
133
|
|
|
129
|
-
def _enforce_strict_provider(
|
|
134
|
+
def _enforce_strict_provider(
|
|
135
|
+
recognizer: TextRecognizer,
|
|
136
|
+
active_provider: str,
|
|
137
|
+
device_id: int = 0,
|
|
138
|
+
) -> None:
|
|
130
139
|
engine = getattr(recognizer, "session", None)
|
|
131
140
|
session = getattr(engine, "session", None)
|
|
132
141
|
if session is None:
|
|
133
142
|
raise RuntimeError("RapidOCR did not expose its ONNX Runtime session")
|
|
134
|
-
|
|
135
|
-
|
|
136
|
-
|
|
137
|
-
|
|
138
|
-
|
|
139
|
-
|
|
140
|
-
disable_fallback = getattr(session, "disable_fallback", None)
|
|
141
|
-
if not callable(disable_fallback):
|
|
142
|
-
raise RuntimeError("RapidOCR ONNX Runtime session cannot disable provider fallback")
|
|
143
|
-
disable_fallback()
|
|
143
|
+
validate_session_provider(
|
|
144
|
+
session,
|
|
145
|
+
active_provider,
|
|
146
|
+
device_id,
|
|
147
|
+
runtime_name="RapidOCR",
|
|
148
|
+
)
|
|
144
149
|
|
|
145
150
|
|
|
146
151
|
def clear_text_recognizer_cache() -> None:
|
|
@@ -0,0 +1,360 @@
|
|
|
1
|
+
# coding: utf-8
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import ctypes
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
import os
|
|
8
|
+
import platform
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
import subprocess
|
|
11
|
+
import sys
|
|
12
|
+
import uuid
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
@dataclass(frozen=True)
|
|
16
|
+
class DeviceIdentity:
|
|
17
|
+
device_id: int | None = None
|
|
18
|
+
name: str = ""
|
|
19
|
+
uuid: str = ""
|
|
20
|
+
verified: bool = False
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
@dataclass(frozen=True)
|
|
24
|
+
class NvidiaDevice:
|
|
25
|
+
index: int
|
|
26
|
+
name: str
|
|
27
|
+
uuid: str
|
|
28
|
+
total_memory_mb: int = 0
|
|
29
|
+
free_memory_mb: int = 0
|
|
30
|
+
driver_version: str = ""
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
@dataclass(frozen=True)
|
|
34
|
+
class DxgiAdapter:
|
|
35
|
+
index: int
|
|
36
|
+
name: str
|
|
37
|
+
luid: str
|
|
38
|
+
dedicated_memory_mb: int = 0
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def resolve_device_identity(provider: str | None, device_id: int | None) -> DeviceIdentity:
|
|
42
|
+
active = str(provider or "")
|
|
43
|
+
if active == "CPUExecutionProvider":
|
|
44
|
+
return DeviceIdentity(name=_cpu_name(), verified=True)
|
|
45
|
+
selected_id = max(0, int(device_id or 0))
|
|
46
|
+
if active in {"CUDAExecutionProvider", "TensorrtExecutionProvider"}:
|
|
47
|
+
device = select_nvidia_device(query_nvidia_devices(), selected_id)
|
|
48
|
+
if device is None:
|
|
49
|
+
return DeviceIdentity(device_id=selected_id)
|
|
50
|
+
return DeviceIdentity(device_id=selected_id, name=device.name, uuid=device.uuid)
|
|
51
|
+
if active == "DmlExecutionProvider":
|
|
52
|
+
adapters = query_dxgi_adapters()
|
|
53
|
+
adapter = next((item for item in adapters if item.index == selected_id), None)
|
|
54
|
+
if adapter is None:
|
|
55
|
+
return DeviceIdentity(device_id=selected_id)
|
|
56
|
+
return DeviceIdentity(device_id=selected_id, name=adapter.name, uuid=adapter.luid)
|
|
57
|
+
return DeviceIdentity()
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def confirm_device_identity(provider: str | None, device_id: int | None) -> DeviceIdentity:
|
|
61
|
+
active = str(provider or "")
|
|
62
|
+
selected = resolve_device_identity(active, device_id)
|
|
63
|
+
if active in {"CUDAExecutionProvider", "TensorrtExecutionProvider"}:
|
|
64
|
+
active_uuids = query_nvidia_process_gpu_uuids()
|
|
65
|
+
if not active_uuids:
|
|
66
|
+
return selected
|
|
67
|
+
devices = query_nvidia_devices()
|
|
68
|
+
matches = [item for item in devices if item.uuid.lower() in active_uuids]
|
|
69
|
+
if len(matches) == 1:
|
|
70
|
+
device = matches[0]
|
|
71
|
+
return DeviceIdentity(
|
|
72
|
+
device_id=max(0, int(device_id or 0)),
|
|
73
|
+
name=device.name,
|
|
74
|
+
uuid=device.uuid,
|
|
75
|
+
verified=True,
|
|
76
|
+
)
|
|
77
|
+
if selected.uuid and selected.uuid.lower() in active_uuids:
|
|
78
|
+
return DeviceIdentity(
|
|
79
|
+
device_id=selected.device_id,
|
|
80
|
+
name=selected.name,
|
|
81
|
+
uuid=selected.uuid,
|
|
82
|
+
verified=True,
|
|
83
|
+
)
|
|
84
|
+
return DeviceIdentity(device_id=max(0, int(device_id or 0)))
|
|
85
|
+
if active == "DmlExecutionProvider":
|
|
86
|
+
return DeviceIdentity(
|
|
87
|
+
device_id=selected.device_id,
|
|
88
|
+
name=selected.name,
|
|
89
|
+
uuid=selected.uuid,
|
|
90
|
+
verified=bool(selected.name),
|
|
91
|
+
)
|
|
92
|
+
return selected
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def query_nvidia_devices() -> tuple[NvidiaDevice, ...]:
|
|
96
|
+
try:
|
|
97
|
+
proc = subprocess.run(
|
|
98
|
+
[
|
|
99
|
+
"nvidia-smi",
|
|
100
|
+
"--query-gpu=index,name,uuid,memory.total,memory.free,driver_version",
|
|
101
|
+
"--format=csv,noheader,nounits",
|
|
102
|
+
],
|
|
103
|
+
check=False,
|
|
104
|
+
capture_output=True,
|
|
105
|
+
text=True,
|
|
106
|
+
encoding="utf-8",
|
|
107
|
+
errors="replace",
|
|
108
|
+
timeout=3.0,
|
|
109
|
+
)
|
|
110
|
+
except Exception:
|
|
111
|
+
return ()
|
|
112
|
+
if proc.returncode != 0:
|
|
113
|
+
return ()
|
|
114
|
+
devices: list[NvidiaDevice] = []
|
|
115
|
+
for raw in proc.stdout.splitlines():
|
|
116
|
+
parts = [part.strip() for part in raw.split(",")]
|
|
117
|
+
if len(parts) < 6:
|
|
118
|
+
continue
|
|
119
|
+
try:
|
|
120
|
+
index = int(parts[0])
|
|
121
|
+
except (TypeError, ValueError):
|
|
122
|
+
continue
|
|
123
|
+
devices.append(
|
|
124
|
+
NvidiaDevice(
|
|
125
|
+
index=index,
|
|
126
|
+
name=parts[1],
|
|
127
|
+
uuid=parts[2],
|
|
128
|
+
total_memory_mb=_safe_int(parts[3]),
|
|
129
|
+
free_memory_mb=_safe_int(parts[4]),
|
|
130
|
+
driver_version=parts[5],
|
|
131
|
+
)
|
|
132
|
+
)
|
|
133
|
+
return tuple(devices)
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def query_nvidia_process_gpu_uuids(pid: int | None = None) -> frozenset[str]:
|
|
137
|
+
expected_pid = int(pid or os.getpid())
|
|
138
|
+
try:
|
|
139
|
+
proc = subprocess.run(
|
|
140
|
+
[
|
|
141
|
+
"nvidia-smi",
|
|
142
|
+
"--query-compute-apps=pid,gpu_uuid",
|
|
143
|
+
"--format=csv,noheader,nounits",
|
|
144
|
+
],
|
|
145
|
+
check=False,
|
|
146
|
+
capture_output=True,
|
|
147
|
+
text=True,
|
|
148
|
+
encoding="utf-8",
|
|
149
|
+
errors="replace",
|
|
150
|
+
timeout=3.0,
|
|
151
|
+
)
|
|
152
|
+
except Exception:
|
|
153
|
+
return frozenset()
|
|
154
|
+
if proc.returncode != 0:
|
|
155
|
+
return frozenset()
|
|
156
|
+
matches: set[str] = set()
|
|
157
|
+
for raw in proc.stdout.splitlines():
|
|
158
|
+
parts = [part.strip() for part in raw.split(",", 1)]
|
|
159
|
+
if len(parts) != 2:
|
|
160
|
+
continue
|
|
161
|
+
try:
|
|
162
|
+
process_id = int(parts[0])
|
|
163
|
+
except (TypeError, ValueError):
|
|
164
|
+
continue
|
|
165
|
+
if process_id == expected_pid and parts[1]:
|
|
166
|
+
matches.add(parts[1].lower())
|
|
167
|
+
return frozenset(matches)
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
def select_nvidia_device(
|
|
171
|
+
devices: tuple[NvidiaDevice, ...],
|
|
172
|
+
logical_device_id: int,
|
|
173
|
+
visible_devices: str | None = None,
|
|
174
|
+
) -> NvidiaDevice | None:
|
|
175
|
+
if logical_device_id < 0:
|
|
176
|
+
return None
|
|
177
|
+
raw_visible = os.environ.get("CUDA_VISIBLE_DEVICES", "") if visible_devices is None else visible_devices
|
|
178
|
+
tokens = [token.strip() for token in raw_visible.split(",") if token.strip()]
|
|
179
|
+
if tokens:
|
|
180
|
+
if logical_device_id >= len(tokens):
|
|
181
|
+
return None
|
|
182
|
+
token = tokens[logical_device_id]
|
|
183
|
+
if token == "-1":
|
|
184
|
+
return None
|
|
185
|
+
if token.isdigit():
|
|
186
|
+
physical_index = int(token)
|
|
187
|
+
return next((item for item in devices if item.index == physical_index), None)
|
|
188
|
+
normalized = token.lower()
|
|
189
|
+
return next(
|
|
190
|
+
(
|
|
191
|
+
item
|
|
192
|
+
for item in devices
|
|
193
|
+
if item.uuid.lower() == normalized or item.uuid.lower().startswith(normalized)
|
|
194
|
+
),
|
|
195
|
+
None,
|
|
196
|
+
)
|
|
197
|
+
return next((item for item in devices if item.index == logical_device_id), None)
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
def query_dxgi_adapters() -> tuple[DxgiAdapter, ...]:
|
|
201
|
+
if os.name != "nt":
|
|
202
|
+
return ()
|
|
203
|
+
try:
|
|
204
|
+
return _query_dxgi_adapters_windows()
|
|
205
|
+
except Exception:
|
|
206
|
+
return ()
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def _query_dxgi_adapters_windows() -> tuple[DxgiAdapter, ...]:
|
|
210
|
+
from ctypes import wintypes
|
|
211
|
+
|
|
212
|
+
class _Guid(ctypes.Structure):
|
|
213
|
+
_fields_ = [
|
|
214
|
+
("Data1", wintypes.DWORD),
|
|
215
|
+
("Data2", wintypes.WORD),
|
|
216
|
+
("Data3", wintypes.WORD),
|
|
217
|
+
("Data4", ctypes.c_ubyte * 8),
|
|
218
|
+
]
|
|
219
|
+
|
|
220
|
+
@classmethod
|
|
221
|
+
def parse(cls, value: str):
|
|
222
|
+
return cls.from_buffer_copy(uuid.UUID(value).bytes_le)
|
|
223
|
+
|
|
224
|
+
class _Luid(ctypes.Structure):
|
|
225
|
+
_fields_ = [("LowPart", wintypes.DWORD), ("HighPart", wintypes.LONG)]
|
|
226
|
+
|
|
227
|
+
class _AdapterDesc1(ctypes.Structure):
|
|
228
|
+
_fields_ = [
|
|
229
|
+
("Description", ctypes.c_wchar * 128),
|
|
230
|
+
("VendorId", wintypes.UINT),
|
|
231
|
+
("DeviceId", wintypes.UINT),
|
|
232
|
+
("SubSysId", wintypes.UINT),
|
|
233
|
+
("Revision", wintypes.UINT),
|
|
234
|
+
("DedicatedVideoMemory", ctypes.c_size_t),
|
|
235
|
+
("DedicatedSystemMemory", ctypes.c_size_t),
|
|
236
|
+
("SharedSystemMemory", ctypes.c_size_t),
|
|
237
|
+
("AdapterLuid", _Luid),
|
|
238
|
+
("Flags", wintypes.UINT),
|
|
239
|
+
]
|
|
240
|
+
|
|
241
|
+
factory = ctypes.c_void_p()
|
|
242
|
+
create_factory = ctypes.windll.dxgi.CreateDXGIFactory1
|
|
243
|
+
create_factory.argtypes = [ctypes.POINTER(_Guid), ctypes.POINTER(ctypes.c_void_p)]
|
|
244
|
+
create_factory.restype = ctypes.c_long
|
|
245
|
+
iid_factory1 = _Guid.parse("770aae78-f26f-4dba-a829-253c83d1b387")
|
|
246
|
+
result = create_factory(ctypes.byref(iid_factory1), ctypes.byref(factory))
|
|
247
|
+
if result < 0 or not factory.value:
|
|
248
|
+
return ()
|
|
249
|
+
|
|
250
|
+
def _method(pointer: ctypes.c_void_p, index: int, prototype):
|
|
251
|
+
table = ctypes.cast(pointer, ctypes.POINTER(ctypes.POINTER(ctypes.c_void_p))).contents
|
|
252
|
+
return prototype(table[index])
|
|
253
|
+
|
|
254
|
+
release_proto = ctypes.WINFUNCTYPE(wintypes.ULONG, ctypes.c_void_p)
|
|
255
|
+
enum_proto = ctypes.WINFUNCTYPE(
|
|
256
|
+
ctypes.c_long,
|
|
257
|
+
ctypes.c_void_p,
|
|
258
|
+
wintypes.UINT,
|
|
259
|
+
ctypes.POINTER(ctypes.c_void_p),
|
|
260
|
+
)
|
|
261
|
+
get_desc_proto = ctypes.WINFUNCTYPE(
|
|
262
|
+
ctypes.c_long,
|
|
263
|
+
ctypes.c_void_p,
|
|
264
|
+
ctypes.POINTER(_AdapterDesc1),
|
|
265
|
+
)
|
|
266
|
+
release_factory = _method(factory, 2, release_proto)
|
|
267
|
+
enum_adapters = _method(factory, 12, enum_proto)
|
|
268
|
+
adapters: list[DxgiAdapter] = []
|
|
269
|
+
try:
|
|
270
|
+
index = 0
|
|
271
|
+
while True:
|
|
272
|
+
adapter = ctypes.c_void_p()
|
|
273
|
+
result = enum_adapters(factory, index, ctypes.byref(adapter))
|
|
274
|
+
if result != 0 or not adapter.value:
|
|
275
|
+
break
|
|
276
|
+
release_adapter = _method(adapter, 2, release_proto)
|
|
277
|
+
try:
|
|
278
|
+
desc = _AdapterDesc1()
|
|
279
|
+
get_desc = _method(adapter, 10, get_desc_proto)
|
|
280
|
+
if get_desc(adapter, ctypes.byref(desc)) >= 0:
|
|
281
|
+
high = int(desc.AdapterLuid.HighPart) & 0xFFFFFFFF
|
|
282
|
+
low = int(desc.AdapterLuid.LowPart) & 0xFFFFFFFF
|
|
283
|
+
adapters.append(
|
|
284
|
+
DxgiAdapter(
|
|
285
|
+
index=index,
|
|
286
|
+
name=str(desc.Description).strip(),
|
|
287
|
+
luid=f"LUID-{high:08X}-{low:08X}",
|
|
288
|
+
dedicated_memory_mb=int(desc.DedicatedVideoMemory // (1024 * 1024)),
|
|
289
|
+
)
|
|
290
|
+
)
|
|
291
|
+
finally:
|
|
292
|
+
release_adapter(adapter)
|
|
293
|
+
index += 1
|
|
294
|
+
finally:
|
|
295
|
+
release_factory(factory)
|
|
296
|
+
return tuple(adapters)
|
|
297
|
+
|
|
298
|
+
|
|
299
|
+
def _cpu_name() -> str:
|
|
300
|
+
if os.name == "nt":
|
|
301
|
+
try:
|
|
302
|
+
import winreg
|
|
303
|
+
|
|
304
|
+
with winreg.OpenKey(
|
|
305
|
+
winreg.HKEY_LOCAL_MACHINE,
|
|
306
|
+
r"HARDWARE\DESCRIPTION\System\CentralProcessor\0",
|
|
307
|
+
) as key:
|
|
308
|
+
value, _ = winreg.QueryValueEx(key, "ProcessorNameString")
|
|
309
|
+
name = str(value or "").strip()
|
|
310
|
+
if name:
|
|
311
|
+
return name
|
|
312
|
+
except Exception:
|
|
313
|
+
pass
|
|
314
|
+
name = str(platform.processor() or "").strip()
|
|
315
|
+
if name:
|
|
316
|
+
return name
|
|
317
|
+
if os.name == "nt":
|
|
318
|
+
return str(os.environ.get("PROCESSOR_IDENTIFIER", "") or "").strip()
|
|
319
|
+
if sys.platform == "darwin":
|
|
320
|
+
try:
|
|
321
|
+
proc = subprocess.run(
|
|
322
|
+
["sysctl", "-n", "machdep.cpu.brand_string"],
|
|
323
|
+
check=False,
|
|
324
|
+
capture_output=True,
|
|
325
|
+
text=True,
|
|
326
|
+
encoding="utf-8",
|
|
327
|
+
errors="replace",
|
|
328
|
+
timeout=2.0,
|
|
329
|
+
)
|
|
330
|
+
if proc.returncode == 0 and proc.stdout.strip():
|
|
331
|
+
return proc.stdout.strip()
|
|
332
|
+
except Exception:
|
|
333
|
+
return ""
|
|
334
|
+
try:
|
|
335
|
+
for raw in Path("/proc/cpuinfo").read_text(encoding="utf-8", errors="ignore").splitlines():
|
|
336
|
+
if raw.lower().startswith("model name"):
|
|
337
|
+
return raw.split(":", 1)[-1].strip()
|
|
338
|
+
except Exception:
|
|
339
|
+
pass
|
|
340
|
+
return ""
|
|
341
|
+
|
|
342
|
+
|
|
343
|
+
def _safe_int(value: str) -> int:
|
|
344
|
+
try:
|
|
345
|
+
return int(float(str(value).strip().replace(",", "")))
|
|
346
|
+
except Exception:
|
|
347
|
+
return 0
|
|
348
|
+
|
|
349
|
+
|
|
350
|
+
__all__ = [
|
|
351
|
+
"DeviceIdentity",
|
|
352
|
+
"DxgiAdapter",
|
|
353
|
+
"NvidiaDevice",
|
|
354
|
+
"confirm_device_identity",
|
|
355
|
+
"query_dxgi_adapters",
|
|
356
|
+
"query_nvidia_devices",
|
|
357
|
+
"query_nvidia_process_gpu_uuids",
|
|
358
|
+
"resolve_device_identity",
|
|
359
|
+
"select_nvidia_device",
|
|
360
|
+
]
|
|
@@ -10,6 +10,7 @@ import os
|
|
|
10
10
|
import subprocess
|
|
11
11
|
import sys
|
|
12
12
|
|
|
13
|
+
from .devices import query_dxgi_adapters, query_nvidia_devices, select_nvidia_device
|
|
13
14
|
from .providers import GPU_PROVIDER_NAMES, ProviderInfo
|
|
14
15
|
|
|
15
16
|
|
|
@@ -24,11 +25,25 @@ class HardwareInfo:
|
|
|
24
25
|
gpu_driver_version: str = ""
|
|
25
26
|
|
|
26
27
|
|
|
27
|
-
@lru_cache(maxsize=
|
|
28
|
-
def detect_hardware_info() -> HardwareInfo:
|
|
28
|
+
@lru_cache(maxsize=8)
|
|
29
|
+
def detect_hardware_info(provider_info: ProviderInfo | None = None) -> HardwareInfo:
|
|
29
30
|
total_mb, free_mb = _memory_status()
|
|
30
|
-
gpu_name
|
|
31
|
-
|
|
31
|
+
gpu_name = ""
|
|
32
|
+
gpu_total_mb = 0
|
|
33
|
+
gpu_free_mb = 0
|
|
34
|
+
gpu_driver = ""
|
|
35
|
+
active = str(getattr(provider_info, "active_provider", "") or "")
|
|
36
|
+
device_id = int(getattr(provider_info, "device_id", 0) or 0)
|
|
37
|
+
if active in {"CUDAExecutionProvider", "TensorrtExecutionProvider"}:
|
|
38
|
+
gpu_name, gpu_total_mb, gpu_free_mb, gpu_driver = _query_nvidia_smi(device_id)
|
|
39
|
+
elif active == "DmlExecutionProvider":
|
|
40
|
+
adapter = next((item for item in query_dxgi_adapters() if item.index == device_id), None)
|
|
41
|
+
if adapter is not None:
|
|
42
|
+
gpu_name = adapter.name
|
|
43
|
+
gpu_total_mb = adapter.dedicated_memory_mb
|
|
44
|
+
elif provider_info is None:
|
|
45
|
+
gpu_name, gpu_total_mb, gpu_free_mb, gpu_driver = _query_nvidia_smi(0)
|
|
46
|
+
if provider_info is None and not gpu_name:
|
|
32
47
|
gpu_name, gpu_total_mb, gpu_driver = _query_windows_video_controller()
|
|
33
48
|
return HardwareInfo(
|
|
34
49
|
logical_processors=max(1, int(os.cpu_count() or 1)),
|
|
@@ -148,30 +163,16 @@ def choose_rec_batch_num(
|
|
|
148
163
|
return 4
|
|
149
164
|
|
|
150
165
|
|
|
151
|
-
def _query_nvidia_smi() -> tuple[str, int, int, str]:
|
|
152
|
-
|
|
153
|
-
|
|
154
|
-
[
|
|
155
|
-
"nvidia-smi",
|
|
156
|
-
"--query-gpu=name,memory.total,memory.free,driver_version",
|
|
157
|
-
"--format=csv,noheader,nounits",
|
|
158
|
-
],
|
|
159
|
-
check=False,
|
|
160
|
-
capture_output=True,
|
|
161
|
-
text=True,
|
|
162
|
-
encoding="utf-8",
|
|
163
|
-
errors="replace",
|
|
164
|
-
timeout=2.0,
|
|
165
|
-
)
|
|
166
|
-
except Exception:
|
|
166
|
+
def _query_nvidia_smi(device_id: int = 0) -> tuple[str, int, int, str]:
|
|
167
|
+
device = select_nvidia_device(query_nvidia_devices(), device_id)
|
|
168
|
+
if device is None:
|
|
167
169
|
return "", 0, 0, ""
|
|
168
|
-
|
|
169
|
-
|
|
170
|
-
|
|
171
|
-
|
|
172
|
-
|
|
173
|
-
|
|
174
|
-
return parts[0], _safe_int(parts[1]), _safe_int(parts[2]), parts[3]
|
|
170
|
+
return (
|
|
171
|
+
device.name,
|
|
172
|
+
device.total_memory_mb,
|
|
173
|
+
device.free_memory_mb,
|
|
174
|
+
device.driver_version,
|
|
175
|
+
)
|
|
175
176
|
|
|
176
177
|
|
|
177
178
|
def _query_windows_video_controller() -> tuple[str, int, str]:
|
|
@@ -3,8 +3,9 @@
|
|
|
3
3
|
from __future__ import annotations
|
|
4
4
|
|
|
5
5
|
import importlib
|
|
6
|
-
from dataclasses import dataclass
|
|
6
|
+
from dataclasses import dataclass, replace
|
|
7
7
|
|
|
8
|
+
from .devices import confirm_device_identity, resolve_device_identity
|
|
8
9
|
from .errors import ProviderError
|
|
9
10
|
|
|
10
11
|
|
|
@@ -22,6 +23,10 @@ class ProviderInfo:
|
|
|
22
23
|
device: str
|
|
23
24
|
gpu_requested: bool
|
|
24
25
|
gpu_runtime_ok: bool
|
|
26
|
+
device_id: int | None = None
|
|
27
|
+
device_name: str = ""
|
|
28
|
+
device_uuid: str = ""
|
|
29
|
+
device_verified: bool = False
|
|
25
30
|
|
|
26
31
|
@property
|
|
27
32
|
def use_cuda(self) -> bool:
|
|
@@ -59,7 +64,7 @@ def detect_providers(prefer: str = "auto") -> ProviderInfo:
|
|
|
59
64
|
if prefer_norm == "cpu":
|
|
60
65
|
if "CPUExecutionProvider" not in available:
|
|
61
66
|
raise ProviderError(f"CPUExecutionProvider unavailable: {available}")
|
|
62
|
-
return
|
|
67
|
+
return _provider_info(
|
|
63
68
|
available_providers=available,
|
|
64
69
|
active_provider="CPUExecutionProvider",
|
|
65
70
|
device="cpu",
|
|
@@ -67,7 +72,7 @@ def detect_providers(prefer: str = "auto") -> ProviderInfo:
|
|
|
67
72
|
gpu_runtime_ok=False,
|
|
68
73
|
)
|
|
69
74
|
if gpu_visible:
|
|
70
|
-
return
|
|
75
|
+
return _provider_info(
|
|
71
76
|
available_providers=available,
|
|
72
77
|
active_provider=gpu_candidates[0],
|
|
73
78
|
device="gpu",
|
|
@@ -80,10 +85,44 @@ def detect_providers(prefer: str = "auto") -> ProviderInfo:
|
|
|
80
85
|
)
|
|
81
86
|
if "CPUExecutionProvider" not in available:
|
|
82
87
|
raise ProviderError(f"no supported ONNX execution provider is available: {available}")
|
|
83
|
-
return
|
|
88
|
+
return _provider_info(
|
|
84
89
|
available_providers=available,
|
|
85
90
|
active_provider="CPUExecutionProvider",
|
|
86
91
|
device="cpu",
|
|
87
92
|
gpu_requested=False,
|
|
88
93
|
gpu_runtime_ok=False,
|
|
89
94
|
)
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def _provider_info(
|
|
98
|
+
*,
|
|
99
|
+
available_providers: tuple[str, ...],
|
|
100
|
+
active_provider: str,
|
|
101
|
+
device: str,
|
|
102
|
+
gpu_requested: bool,
|
|
103
|
+
gpu_runtime_ok: bool,
|
|
104
|
+
) -> ProviderInfo:
|
|
105
|
+
device_id = 0 if active_provider in GPU_PROVIDER_NAMES else None
|
|
106
|
+
identity = resolve_device_identity(active_provider, device_id)
|
|
107
|
+
return ProviderInfo(
|
|
108
|
+
available_providers=available_providers,
|
|
109
|
+
active_provider=active_provider,
|
|
110
|
+
device=device,
|
|
111
|
+
gpu_requested=gpu_requested,
|
|
112
|
+
gpu_runtime_ok=gpu_runtime_ok,
|
|
113
|
+
device_id=identity.device_id,
|
|
114
|
+
device_name=identity.name,
|
|
115
|
+
device_uuid=identity.uuid,
|
|
116
|
+
device_verified=identity.verified,
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def confirm_provider_device(provider_info: ProviderInfo) -> ProviderInfo:
|
|
121
|
+
identity = confirm_device_identity(provider_info.active_provider, provider_info.device_id)
|
|
122
|
+
return replace(
|
|
123
|
+
provider_info,
|
|
124
|
+
device_id=identity.device_id,
|
|
125
|
+
device_name=identity.name,
|
|
126
|
+
device_uuid=identity.uuid,
|
|
127
|
+
device_verified=identity.verified,
|
|
128
|
+
)
|
|
@@ -56,7 +56,7 @@ from .profiles import (
|
|
|
56
56
|
TEXT_DETECTOR_ID,
|
|
57
57
|
TEXT_RECOGNIZER_ID,
|
|
58
58
|
)
|
|
59
|
-
from .providers import ProviderInfo
|
|
59
|
+
from .providers import ProviderInfo, confirm_provider_device
|
|
60
60
|
from .results import Box4P, FormulaRecognitionResult, MathCraftBlock, MixedRecognitionResult, OCRRegion
|
|
61
61
|
|
|
62
62
|
|
|
@@ -618,13 +618,16 @@ class MathCraftRuntime:
|
|
|
618
618
|
except Exception as repair_exc:
|
|
619
619
|
exc = repair_exc
|
|
620
620
|
component_statuses.append(WarmupComponentStatus(model_id=model_id, ready=False, detail=str(exc)))
|
|
621
|
+
provider_info = report.provider_info
|
|
622
|
+
if any(item.ready for item in component_statuses):
|
|
623
|
+
provider_info = confirm_provider_device(provider_info)
|
|
621
624
|
plan = WarmupPlan(
|
|
622
625
|
profile=profile,
|
|
623
626
|
required_models=model_ids,
|
|
624
627
|
missing_models=tuple(missing),
|
|
625
628
|
unsupported_models=tuple(unsupported),
|
|
626
629
|
component_statuses=tuple(component_statuses),
|
|
627
|
-
provider_info=
|
|
630
|
+
provider_info=provider_info,
|
|
628
631
|
ready=not missing
|
|
629
632
|
and not unsupported
|
|
630
633
|
and all(item.ready for item in component_statuses),
|
|
@@ -639,13 +642,15 @@ class MathCraftRuntime:
|
|
|
639
642
|
(
|
|
640
643
|
str(provider_info.device or ""),
|
|
641
644
|
str(provider_info.active_provider or ""),
|
|
645
|
+
str(provider_info.device_id),
|
|
646
|
+
str(provider_info.device_uuid or ""),
|
|
642
647
|
str(provider_info.gpu_runtime_ok),
|
|
643
648
|
)
|
|
644
649
|
)
|
|
645
650
|
cached = self._rec_batch_cache.get(key)
|
|
646
651
|
if cached:
|
|
647
652
|
return cached
|
|
648
|
-
batch = choose_rec_batch_num(provider_info, detect_hardware_info())
|
|
653
|
+
batch = choose_rec_batch_num(provider_info, detect_hardware_info(provider_info))
|
|
649
654
|
self._rec_batch_cache[key] = batch
|
|
650
655
|
return batch
|
|
651
656
|
|
|
@@ -69,6 +69,10 @@ def provider_info_to_json(provider_info) -> dict:
|
|
|
69
69
|
"available_providers": list(provider_info.available_providers),
|
|
70
70
|
"active_provider": provider_info.active_provider,
|
|
71
71
|
"device": provider_info.device,
|
|
72
|
+
"device_id": provider_info.device_id,
|
|
73
|
+
"device_name": provider_info.device_name,
|
|
74
|
+
"device_uuid": provider_info.device_uuid,
|
|
75
|
+
"device_verified": provider_info.device_verified,
|
|
72
76
|
"gpu_requested": provider_info.gpu_requested,
|
|
73
77
|
"gpu_runtime_ok": provider_info.gpu_runtime_ok,
|
|
74
78
|
"use_cuda": provider_info.use_cuda,
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: mathcraft-ocr
|
|
3
|
-
Version: 0.2.
|
|
3
|
+
Version: 0.2.8
|
|
4
4
|
Summary: ONNX-only OCR runtime for mathematical documents
|
|
5
5
|
Author: SakuraMathcraft
|
|
6
6
|
License-Expression: GPL-3.0-only
|
|
@@ -194,7 +194,7 @@ Model artifacts are downloaded from the MathCraft-Models release assets declared
|
|
|
194
194
|
- `cpu`: force CPU.
|
|
195
195
|
- `gpu`: request CUDA-capable ONNX Runtime.
|
|
196
196
|
|
|
197
|
-
The actual provider is available on results through the `provider` field.
|
|
197
|
+
The actual provider is available on recognition results through the `provider` field. Doctor and warmup reports also expose `device_id`, `device_name`, `device_uuid`, and `device_verified` under `provider_info`. GPU sessions bind the reported `device_id` explicitly; `device_verified` becomes true only after the runtime confirms the device used by initialized inference sessions.
|
|
198
198
|
|
|
199
199
|
## Development
|
|
200
200
|
|
|
@@ -1,54 +0,0 @@
|
|
|
1
|
-
# coding: utf-8
|
|
2
|
-
|
|
3
|
-
from __future__ import annotations
|
|
4
|
-
|
|
5
|
-
import importlib
|
|
6
|
-
from functools import lru_cache
|
|
7
|
-
from pathlib import Path
|
|
8
|
-
|
|
9
|
-
from ..providers import GPU_PROVIDER_NAMES, ProviderInfo
|
|
10
|
-
|
|
11
|
-
|
|
12
|
-
def _ort():
|
|
13
|
-
return importlib.import_module("onnxruntime")
|
|
14
|
-
|
|
15
|
-
|
|
16
|
-
def session_providers(provider_info: ProviderInfo) -> list[str]:
|
|
17
|
-
available = list(provider_info.available_providers)
|
|
18
|
-
active = provider_info.active_provider
|
|
19
|
-
if active and active in GPU_PROVIDER_NAMES and "CPUExecutionProvider" in available:
|
|
20
|
-
return [active, "CPUExecutionProvider"]
|
|
21
|
-
if "CPUExecutionProvider" in available:
|
|
22
|
-
return ["CPUExecutionProvider"]
|
|
23
|
-
return available
|
|
24
|
-
|
|
25
|
-
|
|
26
|
-
def create_session(model_path: str | Path, provider_info: ProviderInfo):
|
|
27
|
-
model_path = str(Path(model_path).resolve())
|
|
28
|
-
providers = tuple(session_providers(provider_info))
|
|
29
|
-
session = _create_session_cached(model_path, providers)
|
|
30
|
-
actual = list(session.get_providers() or [])
|
|
31
|
-
active = provider_info.active_provider
|
|
32
|
-
if active and active in GPU_PROVIDER_NAMES and active not in actual:
|
|
33
|
-
raise RuntimeError(
|
|
34
|
-
f"requested ONNX GPU provider {active}, but session providers are {actual}"
|
|
35
|
-
)
|
|
36
|
-
disable_fallback = getattr(session, "disable_fallback", None)
|
|
37
|
-
if not callable(disable_fallback):
|
|
38
|
-
raise RuntimeError("ONNX Runtime session cannot disable provider fallback")
|
|
39
|
-
disable_fallback()
|
|
40
|
-
return session
|
|
41
|
-
|
|
42
|
-
|
|
43
|
-
@lru_cache(maxsize=16)
|
|
44
|
-
def _create_session_cached(model_path: str, providers: tuple[str, ...]):
|
|
45
|
-
ort = _ort()
|
|
46
|
-
return ort.InferenceSession(
|
|
47
|
-
model_path,
|
|
48
|
-
providers=list(providers),
|
|
49
|
-
enable_fallback=False,
|
|
50
|
-
)
|
|
51
|
-
|
|
52
|
-
|
|
53
|
-
def clear_session_cache() -> None:
|
|
54
|
-
_create_session_cached.cache_clear()
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|