cells2table 0.2.2__tar.gz → 0.3.0__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.
- {cells2table-0.2.2 → cells2table-0.3.0}/PKG-INFO +2 -2
- cells2table-0.3.0/cells2table/cli.py +64 -0
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/datamodels/table.py +12 -18
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/docling.py +3 -3
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/models/PaddlePaddle/cell_detection.py +32 -20
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/models/PaddlePaddle/table_classification.py +11 -4
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/models/runtimes/onnx.py +11 -6
- cells2table-0.3.0/cells2table/models/tasks/base.py +28 -0
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/models/tasks/classification.py +1 -1
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/models/tasks/detection.py +1 -1
- cells2table-0.3.0/cells2table/pipelines/__init__.py +4 -0
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/pipelines/base.py +7 -2
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/pipelines/classification_detection.py +5 -5
- cells2table-0.3.0/cells2table/pipelines/paddlepaddle.py +30 -0
- cells2table-0.3.0/cells2table/utils/__init__.py +0 -0
- {cells2table-0.2.2/src/cells2table/models → cells2table-0.3.0/cells2table/utils}/download.py +13 -9
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/utils/visualize.py +1 -1
- {cells2table-0.2.2 → cells2table-0.3.0}/pyproject.toml +10 -5
- cells2table-0.2.2/src/cells2table/__init__.py +0 -3
- cells2table-0.2.2/src/cells2table/models/tasks/base.py +0 -29
- cells2table-0.2.2/src/cells2table/pipelines/__init__.py +0 -3
- cells2table-0.2.2/src/cells2table/pipelines/paddlepaddle.py +0 -45
- {cells2table-0.2.2 → cells2table-0.3.0}/LICENSE +0 -0
- {cells2table-0.2.2 → cells2table-0.3.0}/README.md +0 -0
- {cells2table-0.2.2/src/cells2table/models → cells2table-0.3.0/cells2table}/__init__.py +0 -0
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/datamodels/__init__.py +0 -0
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/datamodels/bbox.py +0 -0
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/models/PaddlePaddle/__init__.py +0 -0
- {cells2table-0.2.2/src/cells2table/models/runtimes → cells2table-0.3.0/cells2table/models}/__init__.py +0 -0
- {cells2table-0.2.2/src/cells2table/utils → cells2table-0.3.0/cells2table/models/runtimes}/__init__.py +0 -0
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/models/tasks/__init__.py +0 -0
- {cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/py.typed +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: cells2table
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.3.0
|
|
4
4
|
Summary: Table image parsing with cell detection models
|
|
5
5
|
Keywords: docling,plugin
|
|
6
6
|
Author: jspast
|
|
@@ -12,7 +12,7 @@ Classifier: Programming Language :: Python :: 3
|
|
|
12
12
|
Requires-Dist: numpy>=2.2.0,<3.0.0
|
|
13
13
|
Requires-Dist: opencv-python>=4.11.0.86,<5.0.0.0
|
|
14
14
|
Requires-Dist: docling>=2.66.0,<3.0.0 ; extra == 'docling'
|
|
15
|
-
Requires-Dist: huggingface-hub>=0.36.0,<
|
|
15
|
+
Requires-Dist: huggingface-hub>=0.36.0,<2.0.0 ; extra == 'huggingface'
|
|
16
16
|
Requires-Dist: onnxruntime>=1.23.2,<2.0.0 ; extra == 'onnx-cpu'
|
|
17
17
|
Requires-Dist: onnxruntime-gpu>=1.23.2,<2.0.0 ; extra == 'onnx-cuda'
|
|
18
18
|
Requires-Dist: onnxruntime-openvino>=1.23.0,<2.0.0 ; extra == 'onnx-openvino'
|
|
@@ -0,0 +1,64 @@
|
|
|
1
|
+
import argparse
|
|
2
|
+
import logging
|
|
3
|
+
from pathlib import Path
|
|
4
|
+
|
|
5
|
+
import cv2
|
|
6
|
+
|
|
7
|
+
from cells2table.pipelines import DefaultPipeline
|
|
8
|
+
from cells2table.utils.visualize import visualize_table
|
|
9
|
+
|
|
10
|
+
logger = logging.getLogger(__name__)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def download(local_dir: Path | str | None = None) -> None:
|
|
14
|
+
"""Download default pipeline models."""
|
|
15
|
+
|
|
16
|
+
log_format = "%(asctime)s\t%(levelname)s\t%(name)s: %(message)s"
|
|
17
|
+
logging.basicConfig(level=logging.INFO, format=log_format)
|
|
18
|
+
|
|
19
|
+
parser = argparse.ArgumentParser()
|
|
20
|
+
parser.add_argument("--local-dir", type=Path, default=None, help="Path to download models to")
|
|
21
|
+
|
|
22
|
+
args = parser.parse_args()
|
|
23
|
+
|
|
24
|
+
DefaultPipeline.download(local_dir=args.local_dir)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def main() -> None:
|
|
28
|
+
"""Basic CLI program for testing."""
|
|
29
|
+
|
|
30
|
+
log_format = "%(asctime)s\t%(levelname)s\t%(name)s: %(message)s"
|
|
31
|
+
logging.basicConfig(level=logging.INFO, format=log_format)
|
|
32
|
+
|
|
33
|
+
parser = argparse.ArgumentParser(description="Load an image from a given path using OpenCV")
|
|
34
|
+
parser.add_argument("image_path", type=Path, help="Path to the image file")
|
|
35
|
+
parser.add_argument("--models-path", type=Path, default=None, help="Path to downloaded models")
|
|
36
|
+
|
|
37
|
+
args = parser.parse_args()
|
|
38
|
+
|
|
39
|
+
if not args.image_path.exists():
|
|
40
|
+
raise FileNotFoundError(f"File does not exist: {args.image_path}")
|
|
41
|
+
|
|
42
|
+
image = cv2.imread(str(args.image_path))
|
|
43
|
+
|
|
44
|
+
if image is None:
|
|
45
|
+
raise ValueError(f"Failed to load image: {args.image_path}")
|
|
46
|
+
|
|
47
|
+
logger.info("Image loaded successfully from %s", args.image_path)
|
|
48
|
+
logger.debug(
|
|
49
|
+
"Image proprieties: width=%d, height=%d, channels=%d, datatype=%s",
|
|
50
|
+
image.shape[1],
|
|
51
|
+
image.shape[0],
|
|
52
|
+
image.shape[2],
|
|
53
|
+
str(image.dtype),
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
table_pipeline = DefaultPipeline(args.models_path)
|
|
57
|
+
tables = table_pipeline([image])
|
|
58
|
+
|
|
59
|
+
for table in tables:
|
|
60
|
+
visualize_table(image, table) # type: ignore
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
if __name__ == "__main__":
|
|
64
|
+
main()
|
|
@@ -1,10 +1,10 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
3
|
from dataclasses import dataclass, field
|
|
4
|
-
from typing import Iterable
|
|
4
|
+
from typing import Iterable
|
|
5
5
|
|
|
6
|
-
from
|
|
7
|
-
from .
|
|
6
|
+
from cells2table.datamodels import BoundingBox
|
|
7
|
+
from cells2table.models.tasks import DetectionResult
|
|
8
8
|
|
|
9
9
|
|
|
10
10
|
@dataclass
|
|
@@ -22,9 +22,9 @@ class Table:
|
|
|
22
22
|
num_rows: int = 0
|
|
23
23
|
num_cols: int = 0
|
|
24
24
|
|
|
25
|
-
@
|
|
26
|
-
def from_detections(cells_det: Iterable[DetectionResult], tolerance: float = 10) -> Table:
|
|
27
|
-
table =
|
|
25
|
+
@classmethod
|
|
26
|
+
def from_detections(cls, cells_det: Iterable[DetectionResult], tolerance: float = 10) -> Table:
|
|
27
|
+
table = cls()
|
|
28
28
|
|
|
29
29
|
for cell_det in cells_det:
|
|
30
30
|
bbox = BoundingBox.from_array(cell_det.bbox)
|
|
@@ -38,20 +38,14 @@ class Table:
|
|
|
38
38
|
self.compute_rows(tolerance)
|
|
39
39
|
self.compute_cols(tolerance)
|
|
40
40
|
|
|
41
|
-
def sort_cells_by_rows(self
|
|
42
|
-
|
|
43
|
-
cells = self.cells
|
|
41
|
+
def sort_cells_by_rows(self) -> None:
|
|
42
|
+
self.cells = sorted(self.cells, key=lambda cell: cell.bbox.t)
|
|
44
43
|
|
|
45
|
-
|
|
46
|
-
|
|
47
|
-
def sort_cells_by_cols(self, cells: Optional[Iterable[Cell]] = None) -> list[Cell]:
|
|
48
|
-
if cells is None:
|
|
49
|
-
cells = self.cells
|
|
50
|
-
|
|
51
|
-
return sorted(self.cells, key=lambda cell: cell.bbox.l)
|
|
44
|
+
def sort_cells_by_cols(self) -> None:
|
|
45
|
+
self.cells = sorted(self.cells, key=lambda cell: cell.bbox.l)
|
|
52
46
|
|
|
53
47
|
def compute_rows(self, tolerance: float) -> None:
|
|
54
|
-
self.
|
|
48
|
+
self.sort_cells_by_rows()
|
|
55
49
|
|
|
56
50
|
row_y = None
|
|
57
51
|
row_num = 0
|
|
@@ -87,7 +81,7 @@ class Table:
|
|
|
87
81
|
self.num_rows = row_num + 1
|
|
88
82
|
|
|
89
83
|
def compute_cols(self, tolerance: float) -> None:
|
|
90
|
-
self.
|
|
84
|
+
self.sort_cells_by_cols()
|
|
91
85
|
|
|
92
86
|
col_x = None
|
|
93
87
|
col_num = 0
|
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import copy
|
|
2
2
|
from collections.abc import Iterable
|
|
3
3
|
from pathlib import Path
|
|
4
|
-
from typing import ClassVar, Literal,
|
|
4
|
+
from typing import ClassVar, Literal, Sequence, Type
|
|
5
5
|
|
|
6
6
|
import numpy
|
|
7
7
|
|
|
@@ -21,7 +21,7 @@ try:
|
|
|
21
21
|
except ImportError:
|
|
22
22
|
raise ImportError("docling is not installed. Unable to initialize plugin.")
|
|
23
23
|
|
|
24
|
-
from . import DefaultPipeline
|
|
24
|
+
from cells2table.pipelines import DefaultPipeline
|
|
25
25
|
|
|
26
26
|
|
|
27
27
|
def get_tokens(page: Page, table_cluster: Cluster, scale: float) -> list[str]:
|
|
@@ -65,7 +65,7 @@ class CustomDoclingTableStructureModel(BaseTableStructureModel):
|
|
|
65
65
|
def __init__(
|
|
66
66
|
self,
|
|
67
67
|
enabled: bool,
|
|
68
|
-
artifacts_path:
|
|
68
|
+
artifacts_path: Path | None,
|
|
69
69
|
options: CustomDoclingTableStructureOptions,
|
|
70
70
|
accelerator_options: AcceleratorOptions,
|
|
71
71
|
**kwargs,
|
{cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/models/PaddlePaddle/cell_detection.py
RENAMED
|
@@ -4,9 +4,9 @@ from typing import Iterable, Iterator, Sequence
|
|
|
4
4
|
import numpy as np
|
|
5
5
|
from numpy.typing import NDArray
|
|
6
6
|
|
|
7
|
-
from
|
|
8
|
-
from
|
|
9
|
-
from
|
|
7
|
+
from cells2table.models.runtimes.onnx import OnnxModel
|
|
8
|
+
from cells2table.models.tasks import DetectionModel, DetectionResult
|
|
9
|
+
from cells2table.utils.download import DownloadOptions, DownloadPlatform
|
|
10
10
|
|
|
11
11
|
HF_REPO_ID = "jspast/paddlepaddle-table-models-onnx"
|
|
12
12
|
CONFIDENCE_THRESHOLD = 0.5
|
|
@@ -58,22 +58,24 @@ class PaddlePaddleCellDetectionModel(DetectionModel, OnnxModel):
|
|
|
58
58
|
scale_factors: Sequence[tuple[int, int]],
|
|
59
59
|
) -> list[Iterator[DetectionResult]]:
|
|
60
60
|
last_cell_idx = 0
|
|
61
|
-
batch_size = len(pred[1])
|
|
62
|
-
generators = []
|
|
63
61
|
cells = pred[0]
|
|
64
62
|
|
|
65
|
-
|
|
66
|
-
|
|
63
|
+
generators = []
|
|
64
|
+
|
|
65
|
+
for i, count in enumerate(pred[1]):
|
|
66
|
+
c = cells[last_cell_idx : last_cell_idx + count]
|
|
67
67
|
c = c[c[:, 1] > CONFIDENCE_THRESHOLD]
|
|
68
|
+
last_cell_idx += count
|
|
68
69
|
|
|
69
|
-
|
|
70
|
+
if not c.size:
|
|
71
|
+
generators.append(iter([]))
|
|
72
|
+
continue
|
|
70
73
|
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
boxes[:, [1, 3]] *= sx
|
|
74
|
+
sx, sy = scale_factors[i]
|
|
75
|
+
scores = c[:, 0]
|
|
76
|
+
boxes = c[:, 2:]
|
|
77
|
+
boxes[:, [0, 2]] *= sy
|
|
78
|
+
boxes[:, [1, 3]] *= sx
|
|
77
79
|
|
|
78
80
|
generators.append((DetectionResult(box, score) for box, score in zip(boxes, scores)))
|
|
79
81
|
|
|
@@ -82,13 +84,23 @@ class PaddlePaddleCellDetectionModel(DetectionModel, OnnxModel):
|
|
|
82
84
|
|
|
83
85
|
class PaddlePaddleWiredCellDetectionModel(PaddlePaddleCellDetectionModel):
|
|
84
86
|
classes = ["wired"]
|
|
85
|
-
|
|
86
|
-
|
|
87
|
-
)
|
|
87
|
+
|
|
88
|
+
@classmethod
|
|
89
|
+
def get_onnx_path(cls) -> str:
|
|
90
|
+
return "wired_table_cell_det.onnx"
|
|
91
|
+
|
|
92
|
+
@classmethod
|
|
93
|
+
def get_download_options(cls) -> DownloadOptions:
|
|
94
|
+
return DownloadOptions(DownloadPlatform.HUGGINGFACE, HF_REPO_ID, [cls.get_onnx_path()])
|
|
88
95
|
|
|
89
96
|
|
|
90
97
|
class PaddlePaddleWirelessCellDetectionModel(PaddlePaddleCellDetectionModel):
|
|
91
98
|
classes = ["wireless"]
|
|
92
|
-
|
|
93
|
-
|
|
94
|
-
)
|
|
99
|
+
|
|
100
|
+
@classmethod
|
|
101
|
+
def get_onnx_path(cls) -> str:
|
|
102
|
+
return "wireless_table_cell_det.onnx"
|
|
103
|
+
|
|
104
|
+
@classmethod
|
|
105
|
+
def get_download_options(cls) -> DownloadOptions:
|
|
106
|
+
return DownloadOptions(DownloadPlatform.HUGGINGFACE, HF_REPO_ID, [cls.get_onnx_path()])
|
{cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/models/PaddlePaddle/table_classification.py
RENAMED
|
@@ -4,9 +4,9 @@ from typing import Iterable, Sequence
|
|
|
4
4
|
import numpy as np
|
|
5
5
|
from numpy.typing import NDArray
|
|
6
6
|
|
|
7
|
-
from
|
|
8
|
-
from
|
|
9
|
-
from
|
|
7
|
+
from cells2table.models.runtimes.onnx import OnnxModel
|
|
8
|
+
from cells2table.models.tasks import ClassificationModel, ClassificationResult
|
|
9
|
+
from cells2table.utils.download import DownloadOptions, DownloadPlatform
|
|
10
10
|
|
|
11
11
|
HF_REPO_ID = "jspast/paddlepaddle-table-models-onnx"
|
|
12
12
|
|
|
@@ -15,7 +15,14 @@ logger = logging.getLogger(__name__)
|
|
|
15
15
|
|
|
16
16
|
class PaddlePaddleTableClassificationModel(ClassificationModel, OnnxModel):
|
|
17
17
|
classes = ["wired", "wireless"]
|
|
18
|
-
|
|
18
|
+
|
|
19
|
+
@classmethod
|
|
20
|
+
def get_onnx_path(cls) -> str:
|
|
21
|
+
return "table_cls.onnx"
|
|
22
|
+
|
|
23
|
+
@classmethod
|
|
24
|
+
def get_download_options(cls) -> DownloadOptions:
|
|
25
|
+
return DownloadOptions(DownloadPlatform.HUGGINGFACE, HF_REPO_ID, [cls.get_onnx_path()])
|
|
19
26
|
|
|
20
27
|
def __call__(self, input: Iterable[NDArray[np.uint8]]) -> list[ClassificationResult]:
|
|
21
28
|
logger.debug("Started preprocessing")
|
|
@@ -1,13 +1,13 @@
|
|
|
1
|
-
from abc import ABC
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
2
|
from pathlib import Path
|
|
3
|
-
from typing import Iterable
|
|
3
|
+
from typing import Iterable
|
|
4
4
|
|
|
5
5
|
import cv2
|
|
6
6
|
import numpy as np
|
|
7
7
|
import onnxruntime as ort
|
|
8
8
|
from numpy.typing import NDArray
|
|
9
9
|
|
|
10
|
-
from
|
|
10
|
+
from cells2table.models.tasks.base import BaseModel
|
|
11
11
|
|
|
12
12
|
|
|
13
13
|
class OnnxModel(BaseModel, ABC):
|
|
@@ -17,7 +17,12 @@ class OnnxModel(BaseModel, ABC):
|
|
|
17
17
|
mean = np.array([0.485, 0.456, 0.406], dtype=np.float32)
|
|
18
18
|
std = np.array([0.229, 0.224, 0.225], dtype=np.float32)
|
|
19
19
|
|
|
20
|
-
|
|
20
|
+
@classmethod
|
|
21
|
+
@abstractmethod
|
|
22
|
+
def get_onnx_path(self) -> str:
|
|
23
|
+
pass
|
|
24
|
+
|
|
25
|
+
def __init__(self, model_path: Path | str | None = None) -> None:
|
|
21
26
|
self.model_path = self.download() if model_path is None else Path(model_path)
|
|
22
27
|
|
|
23
28
|
providers_priority = [
|
|
@@ -26,10 +31,10 @@ class OnnxModel(BaseModel, ABC):
|
|
|
26
31
|
"OpenVINOExecutionProvider",
|
|
27
32
|
"CPUExecutionProvider",
|
|
28
33
|
]
|
|
29
|
-
available_providers = ort.get_available_providers()
|
|
34
|
+
available_providers = ort.get_available_providers()
|
|
30
35
|
|
|
31
36
|
self.session = ort.InferenceSession(
|
|
32
|
-
self.model_path,
|
|
37
|
+
self.model_path / self.get_onnx_path(),
|
|
33
38
|
providers=[p for p in providers_priority if p in available_providers],
|
|
34
39
|
)
|
|
35
40
|
|
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
from cells2table.utils.download import DownloadOptions
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class BaseModel(ABC):
|
|
9
|
+
"""Base interface for models of any type."""
|
|
10
|
+
|
|
11
|
+
model_path: Path
|
|
12
|
+
|
|
13
|
+
@abstractmethod
|
|
14
|
+
def __init__(self, model_path: Path | str | None = None) -> None:
|
|
15
|
+
pass
|
|
16
|
+
|
|
17
|
+
@abstractmethod
|
|
18
|
+
def __call__(self, input: Any):
|
|
19
|
+
pass
|
|
20
|
+
|
|
21
|
+
@classmethod
|
|
22
|
+
@abstractmethod
|
|
23
|
+
def get_download_options(cls) -> DownloadOptions:
|
|
24
|
+
pass
|
|
25
|
+
|
|
26
|
+
@classmethod
|
|
27
|
+
def download(cls, *, local_dir: Path | str | None = None) -> Path:
|
|
28
|
+
return cls.get_download_options().download(local_dir=local_dir)
|
|
@@ -1,15 +1,20 @@
|
|
|
1
1
|
from abc import ABC, abstractmethod
|
|
2
2
|
from pathlib import Path
|
|
3
|
-
from typing import Any
|
|
3
|
+
from typing import Any
|
|
4
4
|
|
|
5
5
|
|
|
6
6
|
class BasePipeline(ABC):
|
|
7
7
|
"""Base interface for pipelines of any type."""
|
|
8
8
|
|
|
9
9
|
@abstractmethod
|
|
10
|
-
def __init__(self, models_path:
|
|
10
|
+
def __init__(self, models_path: Path | str | None = None) -> None:
|
|
11
11
|
pass
|
|
12
12
|
|
|
13
13
|
@abstractmethod
|
|
14
14
|
def __call__(self, input: Any):
|
|
15
15
|
pass
|
|
16
|
+
|
|
17
|
+
@classmethod
|
|
18
|
+
@abstractmethod
|
|
19
|
+
def download(cls, *, local_dir: Path | str | None = None) -> Path:
|
|
20
|
+
pass
|
{cells2table-0.2.2/src → cells2table-0.3.0}/cells2table/pipelines/classification_detection.py
RENAMED
|
@@ -1,11 +1,11 @@
|
|
|
1
1
|
import logging
|
|
2
2
|
from abc import ABC, abstractmethod
|
|
3
3
|
from pathlib import Path
|
|
4
|
-
from typing import Any, Iterable
|
|
4
|
+
from typing import Any, Iterable
|
|
5
5
|
|
|
6
|
-
from
|
|
7
|
-
from
|
|
8
|
-
from .base import BasePipeline
|
|
6
|
+
from cells2table.datamodels import Table
|
|
7
|
+
from cells2table.models.tasks import ClassificationModel, DetectionModel
|
|
8
|
+
from cells2table.pipelines.base import BasePipeline
|
|
9
9
|
|
|
10
10
|
logger = logging.getLogger(__name__)
|
|
11
11
|
|
|
@@ -17,7 +17,7 @@ class ClassificationDetectionPipeline(BasePipeline, ABC):
|
|
|
17
17
|
detection_models: list[DetectionModel]
|
|
18
18
|
|
|
19
19
|
@abstractmethod
|
|
20
|
-
def __init__(self, models_path:
|
|
20
|
+
def __init__(self, models_path: Path | str | None = None) -> None:
|
|
21
21
|
"""Initialize the models."""
|
|
22
22
|
|
|
23
23
|
pass
|
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
from pathlib import Path
|
|
2
|
+
|
|
3
|
+
from cells2table.models.PaddlePaddle import (
|
|
4
|
+
PaddlePaddleTableClassificationModel,
|
|
5
|
+
PaddlePaddleWiredCellDetectionModel,
|
|
6
|
+
PaddlePaddleWirelessCellDetectionModel,
|
|
7
|
+
)
|
|
8
|
+
from cells2table.pipelines.classification_detection import ClassificationDetectionPipeline
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class PaddlePaddleTablePipeline(ClassificationDetectionPipeline):
|
|
12
|
+
"""A table pipeline combining PaddlePaddle classification and detection models."""
|
|
13
|
+
|
|
14
|
+
def __init__(self, models_path: Path | str | None = None) -> None:
|
|
15
|
+
"""Initialize models from the provided path or download them."""
|
|
16
|
+
|
|
17
|
+
models_path = self.download() if models_path is None else Path(models_path)
|
|
18
|
+
|
|
19
|
+
self.classification_model = PaddlePaddleTableClassificationModel(models_path)
|
|
20
|
+
self.detection_models = [
|
|
21
|
+
PaddlePaddleWiredCellDetectionModel(models_path),
|
|
22
|
+
PaddlePaddleWirelessCellDetectionModel(models_path),
|
|
23
|
+
]
|
|
24
|
+
|
|
25
|
+
@classmethod
|
|
26
|
+
def download(cls, *, local_dir: Path | str | None = None) -> Path:
|
|
27
|
+
PaddlePaddleTableClassificationModel.download(local_dir=local_dir)
|
|
28
|
+
PaddlePaddleWiredCellDetectionModel.download(local_dir=local_dir)
|
|
29
|
+
path = PaddlePaddleWirelessCellDetectionModel.download(local_dir=local_dir)
|
|
30
|
+
return path
|
|
File without changes
|
{cells2table-0.2.2/src/cells2table/models → cells2table-0.3.0/cells2table/utils}/download.py
RENAMED
|
@@ -13,18 +13,22 @@ class DownloadPlatform(Enum):
|
|
|
13
13
|
class DownloadOptions(NamedTuple):
|
|
14
14
|
platform: DownloadPlatform
|
|
15
15
|
repo_id: str
|
|
16
|
-
|
|
16
|
+
files: list[str] | None = None
|
|
17
17
|
|
|
18
|
+
def download(self, *, local_dir: Path | str | None = None) -> Path:
|
|
19
|
+
match self.platform:
|
|
20
|
+
case DownloadPlatform.HUGGINGFACE:
|
|
21
|
+
path = download_hf_model(self.repo_id, files=self.files, local_dir=local_dir)
|
|
18
22
|
|
|
19
|
-
|
|
20
|
-
match options.platform:
|
|
21
|
-
case DownloadPlatform.HUGGINGFACE:
|
|
22
|
-
path = download_hf_model(options.repo_id)
|
|
23
|
+
return path
|
|
23
24
|
|
|
24
|
-
return path / options.model_path
|
|
25
25
|
|
|
26
|
-
|
|
27
|
-
|
|
26
|
+
def download_hf_model(
|
|
27
|
+
repo_id: str,
|
|
28
|
+
*,
|
|
29
|
+
files: list[str] | None = None,
|
|
30
|
+
local_dir: Path | str | None = None,
|
|
31
|
+
) -> Path:
|
|
28
32
|
"""Download a repository from Hugging Face and return its path."""
|
|
29
33
|
|
|
30
34
|
try:
|
|
@@ -36,6 +40,6 @@ def download_hf_model(repo_id: str) -> Path:
|
|
|
36
40
|
disable_progress_bars()
|
|
37
41
|
|
|
38
42
|
logger.info("Downloading HF repo %s", repo_id)
|
|
39
|
-
download_path = snapshot_download(repo_id=repo_id)
|
|
43
|
+
download_path = snapshot_download(repo_id=repo_id, allow_patterns=files, local_dir=local_dir)
|
|
40
44
|
|
|
41
45
|
return Path(download_path)
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "cells2table"
|
|
3
|
-
version = "0.
|
|
3
|
+
version = "0.3.0"
|
|
4
4
|
description = "Table image parsing with cell detection models"
|
|
5
5
|
readme = "README.md"
|
|
6
6
|
requires-python = ">=3.12,<4.0"
|
|
@@ -24,7 +24,8 @@ Homepage = "https://github.com/jspast/cells2table"
|
|
|
24
24
|
Issues = "https://github.com/jspast/cells2table/issues"
|
|
25
25
|
|
|
26
26
|
[project.scripts]
|
|
27
|
-
cells2table = "cli:main"
|
|
27
|
+
cells2table = "cells2table.cli:main"
|
|
28
|
+
cells2table-download = "cells2table.cli:download"
|
|
28
29
|
|
|
29
30
|
[project.entry-points.docling]
|
|
30
31
|
cells2table = "cells2table.docling"
|
|
@@ -34,7 +35,7 @@ docling = [
|
|
|
34
35
|
"docling (>=2.66.0,<3.0.0)",
|
|
35
36
|
]
|
|
36
37
|
huggingface = [
|
|
37
|
-
"huggingface-hub (>=0.36.0,<
|
|
38
|
+
"huggingface-hub (>=0.36.0,<2.0.0)",
|
|
38
39
|
]
|
|
39
40
|
onnx_cpu = [
|
|
40
41
|
"onnxruntime (>=1.23.2,<2.0.0)",
|
|
@@ -49,12 +50,12 @@ onnx_openvino = [
|
|
|
49
50
|
[dependency-groups]
|
|
50
51
|
dev = [
|
|
51
52
|
"cells2table",
|
|
52
|
-
"pytest (>=7.0,<
|
|
53
|
+
"pytest (>=7.0,<10.0)",
|
|
53
54
|
"ruff (>=0.5.0,<1.0.0)",
|
|
54
55
|
]
|
|
55
56
|
|
|
56
57
|
[build-system]
|
|
57
|
-
requires = ["uv_build>=0.9.17,<0.
|
|
58
|
+
requires = ["uv_build>=0.9.17,<1.0.0"]
|
|
58
59
|
build-backend = "uv_build"
|
|
59
60
|
|
|
60
61
|
[tool.ruff]
|
|
@@ -69,3 +70,7 @@ conflicts = [
|
|
|
69
70
|
{ extra = "onnx_openvino" }
|
|
70
71
|
]
|
|
71
72
|
]
|
|
73
|
+
|
|
74
|
+
[tool.uv.build-backend]
|
|
75
|
+
module-name = "cells2table"
|
|
76
|
+
module-root = ""
|
|
@@ -1,29 +0,0 @@
|
|
|
1
|
-
from abc import ABC, abstractmethod
|
|
2
|
-
from pathlib import Path
|
|
3
|
-
from typing import Any, Optional
|
|
4
|
-
|
|
5
|
-
from ..download import DownloadOptions, download
|
|
6
|
-
|
|
7
|
-
|
|
8
|
-
class BaseModel(ABC):
|
|
9
|
-
"""Base interface for models of any type."""
|
|
10
|
-
|
|
11
|
-
model_path: Path
|
|
12
|
-
download_options: Optional[DownloadOptions] = None
|
|
13
|
-
|
|
14
|
-
@abstractmethod
|
|
15
|
-
def __init__(self, model_path: Optional[Path | str] = None) -> None:
|
|
16
|
-
pass
|
|
17
|
-
|
|
18
|
-
@abstractmethod
|
|
19
|
-
def __call__(self, input: Any):
|
|
20
|
-
pass
|
|
21
|
-
|
|
22
|
-
@classmethod
|
|
23
|
-
def download(cls) -> Path:
|
|
24
|
-
if cls.download_options is not None:
|
|
25
|
-
return download(cls.download_options)
|
|
26
|
-
else:
|
|
27
|
-
raise NotImplementedError(
|
|
28
|
-
"Download is not implemented for this model. Please provide a path."
|
|
29
|
-
)
|
|
@@ -1,45 +0,0 @@
|
|
|
1
|
-
import logging
|
|
2
|
-
from pathlib import Path
|
|
3
|
-
from typing import Optional
|
|
4
|
-
|
|
5
|
-
from ..models.PaddlePaddle import (
|
|
6
|
-
PaddlePaddleTableClassificationModel,
|
|
7
|
-
PaddlePaddleWiredCellDetectionModel,
|
|
8
|
-
PaddlePaddleWirelessCellDetectionModel,
|
|
9
|
-
)
|
|
10
|
-
from .classification_detection import ClassificationDetectionPipeline
|
|
11
|
-
|
|
12
|
-
logger = logging.getLogger(__name__)
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
class PaddlePaddleTablePipeline(ClassificationDetectionPipeline):
|
|
16
|
-
"""A table pipeline combining PaddlePaddle classification and detection models."""
|
|
17
|
-
|
|
18
|
-
def __init__(self, models_path: Optional[Path | str] = None) -> None:
|
|
19
|
-
"""Initialize models from the provided path or download them.
|
|
20
|
-
|
|
21
|
-
As the models are all in the same repository, do the download only once.
|
|
22
|
-
"""
|
|
23
|
-
|
|
24
|
-
self.detection_models = []
|
|
25
|
-
cls_path, wired_path, wireless_path = None, None, None
|
|
26
|
-
|
|
27
|
-
if models_path is not None:
|
|
28
|
-
models_path = Path(models_path)
|
|
29
|
-
cls_path = (
|
|
30
|
-
models_path / PaddlePaddleTableClassificationModel.download_options.model_path
|
|
31
|
-
)
|
|
32
|
-
|
|
33
|
-
self.classification_model = PaddlePaddleTableClassificationModel(cls_path)
|
|
34
|
-
|
|
35
|
-
wired_path = (
|
|
36
|
-
self.classification_model.model_path.parent
|
|
37
|
-
/ PaddlePaddleWiredCellDetectionModel.download_options.model_path
|
|
38
|
-
)
|
|
39
|
-
wireless_path = (
|
|
40
|
-
self.classification_model.model_path.parent
|
|
41
|
-
/ PaddlePaddleWirelessCellDetectionModel.download_options.model_path
|
|
42
|
-
)
|
|
43
|
-
|
|
44
|
-
self.detection_models.append(PaddlePaddleWiredCellDetectionModel(wired_path))
|
|
45
|
-
self.detection_models.append(PaddlePaddleWirelessCellDetectionModel(wireless_path))
|
|
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
|