Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 21 additions & 3 deletions extract-core/extract_core/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from icij_common.pydantic_utils import make_enum_discriminator, tagged_union
from pydantic import Discriminator

from .configs import BasePipelineConfig, PipelineType
from .configs import BasePipelineConfig, PipelineType, ResultBufferConfig
from .objects import (
BaseModel,
ConversionOutput,
Expand All @@ -21,9 +21,24 @@
from .pipeline import Pipeline

try:
from .docling_ import DoclingFormatOption, DoclingPipelineConfig
from .docling_ import (
BatchConcurrencySettings,
DoclingFormatOption,
DoclingPipelineConfig,
DoclingSettings,
)
except ModuleNotFoundError:
DoclingPipelineConfig, DoclingFormatOption = None, None
(
BatchConcurrencySettings,
DoclingFormatOption,
DoclingPipelineConfig,
DoclingSettings,
) = (
None,
None,
None,
None,
)

try:
from .marker_ import MarkerPipelineConfig
Expand Down Expand Up @@ -66,4 +81,7 @@
"Result",
"Status",
"SupportedExt",
"ResultBufferConfig",
"DoclingSettings",
"BatchConcurrencySettings",
]
10 changes: 8 additions & 2 deletions extract-core/extract_core/configs.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,13 @@
from abc import ABC, abstractmethod
from enum import StrEnum
from pathlib import Path
from typing import ClassVar

from icij_common.pydantic_utils import icij_config, merge_configs, no_enum_values_config
from icij_common.registrable import RegistrableConfig
from pydantic import Field
from pydantic import ByteSize, Field

from .objects import Device, SupportedExt
from .objects import BaseModel, Device, SupportedExt


class PipelineType(StrEnum):
Expand All @@ -27,3 +28,8 @@ class BasePipelineConfig(RegistrableConfig, ABC):
@classmethod
@abstractmethod
def supported_exts(cls) -> set[SupportedExt]: ...


class ResultBufferConfig(BaseModel):
max_size: ByteSize = "500MiB"
root: Path | None = None
22 changes: 15 additions & 7 deletions extract-core/extract_core/docling_.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import importlib
from functools import cache
from typing import Annotated, Any, ClassVar, TypeVar, get_type_hints
from typing import Annotated, Any, ClassVar, get_type_hints

from docling.datamodel.backend_options import BackendOptions, BaseBackendOptions
from docling.datamodel.base_models import (
Expand All @@ -21,7 +21,9 @@
ThreadedPdfPipelineOptions,
)
from docling.datamodel.settings import (
BatchConcurrencySettings,
BatchConcurrencySettings as DoclingBatchConcurrencySettings,
)
from docling.datamodel.settings import (
DebugSettings,
InferenceSettings,
)
Expand All @@ -40,7 +42,7 @@
)
from pydantic_core.core_schema import SerializerFunctionWrapHandler

from .configs import BasePipelineConfig, PipelineType
from .configs import BasePipelineConfig, PipelineType, ResultBufferConfig
from .objects import BaseModel, Device, SupportedExt
from .utils import all_subclasses

Expand Down Expand Up @@ -69,10 +71,7 @@ def _validate_pipeline_opts(v: PipelineOptions) -> PipelineOptions:
return v


T = TypeVar("T")


def _find_subcls(cls: type[T], name: str) -> type[T]:
def _find_subcls[T](cls: type[T], name: str) -> type[T]:
# Check if the class available
for c in all_subclasses(cls):
if c.__name__ == name:
Expand Down Expand Up @@ -264,6 +263,13 @@ def _default_format_opts() -> dict[InputFormat, DoclingFormatOption]:
}


class BatchConcurrencySettings(DoclingBatchConcurrencySettings):
# process up to 16 pages in || on GPU
page_batch_size: int = 16
# call convert_all with at most page_batch_size * page_batch_size
max_page_batches: int = 2


class DoclingSettings(BaseModel):
perf: BatchConcurrencySettings = Field(default_factory=BatchConcurrencySettings)
debug: DebugSettings = Field(default_factory=DebugSettings)
Expand All @@ -276,7 +282,9 @@ class DoclingPipelineConfig(BasePipelineConfig):
format_options: dict[InputFormat, DoclingFormatOption] = Field(
default_factory=_default_format_opts
)

settings: DoclingSettings = Field(default_factory=DoclingSettings)
result_buffer: ResultBufferConfig = Field(default_factory=ResultBufferConfig)

@classmethod
@cache
Expand Down
50 changes: 29 additions & 21 deletions extract-core/extract_core/objects.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,17 +8,16 @@
from functools import cache
from io import BytesIO
from pathlib import Path
from typing import Annotated, Any, NoReturn, Self
from typing import Any, NoReturn, Self

from docling.datamodel.accelerator_options import AcceleratorDevice
from icij_common.pydantic_utils import (
icij_config,
merge_configs,
no_enum_values_config,
safe_copy,
)
from pydantic import AfterValidator, Field, TypeAdapter
from pydantic import BaseModel as _BaseModel
from pydantic import Field, TypeAdapter

logger = logging.getLogger(__name__)
base_config = merge_configs(icij_config(), no_enum_values_config())
Expand Down Expand Up @@ -120,8 +119,8 @@ def to_marker(self) -> str:

class Status(StrEnum):
FAILURE = "failure"
SUCCESS = "success"
PARTIAL_SUCCESS = "partial_success"
SUCCESS = "success"

@classmethod
def from_docling(cls, v: Any) -> Self:
Expand All @@ -139,6 +138,26 @@ def from_docling(cls, v: Any) -> Self:
def allows_conversion(self) -> bool:
return self is Status.SUCCESS or self is Status.PARTIAL_SUCCESS

def __add__(self, other: "Status") -> "Status":
if not isinstance(other, Status):
msg = (
f"can't add {other} of type {other.__class__.__name__} "
f"to {self.__class__.__name__}"
)
raise TypeError(msg)
statuses = sorted((self, other), key=lambda x: x.value)
match statuses:
case (Status.FAILURE, Status.FAILURE):
return Status.FAILURE
case (Status.FAILURE, Status.SUCCESS):
return Status.PARTIAL_SUCCESS
case (_, Status.PARTIAL_SUCCESS) | (Status.PARTIAL_SUCCESS, _):
return Status.PARTIAL_SUCCESS
case (Status.SUCCESS, Status.SUCCESS):
return Status.SUCCESS
case _:
raise ValueError(f"unexpected value {statuses}")


class Error(BaseModel):
id: str
Expand Down Expand Up @@ -179,30 +198,24 @@ def _id_title(title: str) -> str:
class InputDoc(BaseModel):
ext: SupportedExt
path: Path
content: bytes | None = None
n_pages: int

@classmethod
def from_path(cls, path: str | Path) -> Self:
def from_path(cls, path: str | Path, n_pages: int) -> Self:
if isinstance(path, str):
path = Path(path)
ext = SupportedExt(path.suffix)
return cls(path=path, ext=ext)
return cls(path=path, ext=ext, n_pages=n_pages)

def to_docling(self): # noqa: ANN201
from docling_core.types.io import DocumentStream # noqa: PLC0415

if self.content is not None:
return DocumentStream(name=str(self.path), stream=BytesIO(self.content))

if not self.path.suffix:
return DocumentStream(
name=str(self.path), stream=BytesIO(self.path.read_bytes())
)
return self.path

def without_content(self) -> Self:
return safe_copy(self, update={"content": None})


Ranges = list[tuple[int, int]]

Expand All @@ -223,6 +236,7 @@ def from_pages_bytes_sizes(cls, sizes: Sequence[int]) -> Self:
class ConversionOutput(BaseModel):
path: Path
pages: Pages = Field(default_factory=Pages)
confidence: float | None


class MarkdownDoc(ConversionOutput):
Expand All @@ -235,20 +249,14 @@ def _valid_conversion_statuses(cls) -> set:
return {ConversionStatus.SUCCESS, ConversionStatus.PARTIAL_SUCCESS}


def _input_should_not_have_content(value: InputDoc) -> InputDoc:
if value.content is not None:
raise ValueError(f"response input can't have content, but got {value}")
return value


class _BaseResult(BaseModel, ABC):
input: InputDoc
status: Status
errors: list[Error] = []


class ResponseResult(_BaseResult):
input: Annotated[InputDoc, AfterValidator(func=_input_should_not_have_content)]
input: InputDoc
output_path: Path


Expand All @@ -258,7 +266,7 @@ class Result(_BaseResult):

def to_response(self) -> ResponseResult:
return ResponseResult(
input=self.input.without_content(),
input=self.input,
status=self.status,
errors=self.errors,
output_path=self.output.path,
Expand Down
10 changes: 4 additions & 6 deletions extract-core/extract_core/pipeline.py
Original file line number Diff line number Diff line change
@@ -1,26 +1,24 @@
from abc import ABC, abstractmethod
from collections.abc import AsyncGenerator, Iterable
from collections.abc import AsyncIterable, Iterable
from pathlib import Path
from typing import Generic, Self, TypeVar
from typing import Self

from icij_common.registrable import RegistrableFromConfig

from extract_core import BasePipelineConfig

from .objects import InputDoc, OutputFormat, Result

C = TypeVar("C", bound="BasePipelineConfig")


class Pipeline(RegistrableFromConfig, Generic[C], ABC):
class Pipeline[C: BasePipelineConfig](RegistrableFromConfig, ABC):
def __init__(self, config: C):
self._config = config
self._device = self._config.device

@abstractmethod
async def extract_content(
self, docs: Iterable[InputDoc], output_format: OutputFormat, output_path: Path
) -> AsyncGenerator[Result, None]: ...
) -> AsyncIterable[Result]: ...

@classmethod
def _from_config(cls, config: C) -> Self:
Expand Down
7 changes: 1 addition & 6 deletions extract-core/extract_core/utils.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,4 @@
from typing import TypeVar

T = TypeVar("T")


def all_subclasses(cls: type[T]) -> set[type[T]]:
def all_subclasses[T](cls: type[T]) -> set[type[T]]:
return set(cls.__subclasses__()).union(
[s for c in cls.__subclasses__() for s in all_subclasses(c)]
)
24 changes: 23 additions & 1 deletion extract-core/tests/test_objects.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import pytest
from docling.datamodel.accelerator_options import AcceleratorDevice, AcceleratorOptions
from docling.datamodel.base_models import InputFormat
from docling.datamodel.pipeline_options import (
Expand All @@ -6,7 +7,7 @@
)
from docling.document_converter import PdfFormatOption
from extract_core import DoclingPipelineConfig, PipelineConfig
from extract_core.objects import Device
from extract_core.objects import Device, Status
from pydantic import TypeAdapter


Expand Down Expand Up @@ -43,3 +44,24 @@ def test_docling_pipeline_config() -> None:
)
)
assert pdf_pipeline_options.model_dump() == expected_options.model_dump()


@pytest.mark.parametrize(
("left", "right", "expected_status"),
[
(Status.FAILURE, Status.FAILURE, Status.FAILURE),
(Status.FAILURE, Status.PARTIAL_SUCCESS, Status.PARTIAL_SUCCESS),
(Status.FAILURE, Status.SUCCESS, Status.PARTIAL_SUCCESS),
(Status.PARTIAL_SUCCESS, Status.FAILURE, Status.PARTIAL_SUCCESS),
(Status.PARTIAL_SUCCESS, Status.PARTIAL_SUCCESS, Status.PARTIAL_SUCCESS),
(Status.PARTIAL_SUCCESS, Status.SUCCESS, Status.PARTIAL_SUCCESS),
(Status.SUCCESS, Status.FAILURE, Status.PARTIAL_SUCCESS),
(Status.SUCCESS, Status.PARTIAL_SUCCESS, Status.PARTIAL_SUCCESS),
(Status.SUCCESS, Status.SUCCESS, Status.SUCCESS),
],
)
def test_add_statuses(left: Status, right: Status, expected_status: Status) -> None:
# When
status = left + right
# Then
assert status == expected_status
Loading