diff --git a/docling/.agents/skills/docling/references/cli.md b/docling/.agents/skills/docling/references/cli.md index b76fbe65f6..40e7fa1d98 100644 --- a/docling/.agents/skills/docling/references/cli.md +++ b/docling/.agents/skills/docling/references/cli.md @@ -50,9 +50,19 @@ docling report.pdf --pipeline vlm --output /tmp/ docling report.pdf --pipeline vlm --vlm-model granite_docling --output /tmp/ docling report.pdf --pipeline vlm --vlm-model smoldocling --output /tmp/ docling report.pdf --pipeline vlm --vlm-model nemotron_parse_v2 --output /tmp/ +docling chart.png --from image --to dclx --pipeline vlm --enrich-chart-extraction docling report.pdf --pipeline native --from pdf --output /tmp/ ``` +VLM picture enrichment runs after conversion. `--enrich-chart-extraction` +classifies pictures first and adds chart data only where the VLM has not +already provided it. +Use `--chart-extraction-preset granite_vision` for the older CSV-only +`ibm-granite/granite-vision-3.3-2b-chart2csv-preview` checkpoint. The default +`granite_vision_v4` preset uses Granite Vision 4.1 4B, selecting MLX on +compatible Apple Silicon systems and Transformers elsewhere. To require MLX, +use `--chart-extraction-preset granite_vision_v4_mlx` with the same repository. + For PDFs, visible horizontal and vertical rules are used as reading-order signals by default. Disable this to compare against rule-free ordering: diff --git a/docling/cli/main.py b/docling/cli/main.py index 2d56e36356..517fca5066 100644 --- a/docling/cli/main.py +++ b/docling/cli/main.py @@ -114,6 +114,7 @@ InputFormat, OutputFormat, ) +from docling.datamodel.chart_extraction_options import ChartExtractionVlmEngineOptions from docling.datamodel.document import ConversionResult, DoclingVersion from docling.datamodel.pipeline_options import ( AsrPipelineOptions, @@ -292,6 +293,7 @@ def _expand_from_formats(from_formats: list[str] | None) -> list[InputFormat]: # Get available VLM presets from the registry vlm_preset_ids = VlmConvertOptions.list_preset_ids() +chart_extraction_preset_ids = ChartExtractionVlmEngineOptions.list_preset_ids() DOCLING_ASCII_ART = r""" ████ ██████ @@ -1070,6 +1072,16 @@ def convert( # noqa: C901 ..., help="Enable chart data extraction from bar, pie, and line charts." ), ] = False, + chart_extraction_preset: Annotated[ + str, + typer.Option( + "--chart-extraction-preset", + help=( + "Choose the chart extraction preset. Available presets: " + f"{', '.join(chart_extraction_preset_ids)}" + ), + ), + ] = "granite_vision_v4", artifacts_path: Annotated[ Path | None, typer.Option(..., help="If provided, the location of the model artifacts."), @@ -1483,6 +1495,17 @@ def _resolve_pdf_backend() -> tuple[type[PdfDocumentBackend], PdfBackendOptions] table_structure_factory.create_options(kind=table_structure_engine) ) + try: + chart_extraction_options = ChartExtractionVlmEngineOptions.from_preset( + chart_extraction_preset + ) + except KeyError as exc: + raise typer.BadParameter( + f"Unknown chart extraction preset {chart_extraction_preset!r}. " + f"Available presets: {', '.join(chart_extraction_preset_ids)}", + param_hint="--chart-extraction-preset", + ) from exc + if pipeline in {ProcessingPipeline.STANDARD, ProcessingPipeline.LEGACY}: pipeline_cls = ( LegacyStandardPdfPipeline @@ -1504,6 +1527,7 @@ def _resolve_pdf_backend() -> tuple[type[PdfDocumentBackend], PdfBackendOptions] do_picture_description=enrich_picture_description, do_picture_classification=enrich_picture_classes, do_chart_extraction=enrich_chart_extraction, + chart_extraction_options=chart_extraction_options, document_timeout=document_timeout, ) if isinstance( @@ -1541,6 +1565,7 @@ def _resolve_pdf_backend() -> tuple[type[PdfDocumentBackend], PdfBackendOptions] do_picture_description=enrich_picture_description, do_picture_classification=enrich_picture_classes, do_chart_extraction=enrich_chart_extraction, + chart_extraction_options=chart_extraction_options, ) if artifacts_path is not None: simple_format_option.artifacts_path = artifacts_path @@ -1637,6 +1662,7 @@ def _resolve_pdf_backend() -> tuple[type[PdfDocumentBackend], PdfBackendOptions] do_picture_description=enrich_picture_description, do_picture_classification=enrich_picture_classes, do_chart_extraction=enrich_chart_extraction, + chart_extraction_options=chart_extraction_options, document_timeout=document_timeout, ) if parser_threads is not None: @@ -1671,7 +1697,15 @@ def _resolve_pdf_backend() -> tuple[type[PdfDocumentBackend], PdfBackendOptions] pipeline_options = VlmPipelineOptions( accelerator_options=accelerator_options, enable_remote_services=enable_remote_services, + do_picture_classification=enrich_picture_classes, + do_picture_description=enrich_picture_description, + do_chart_extraction=enrich_chart_extraction, + chart_extraction_options=chart_extraction_options, ) + if picture_description_max_new_tokens is not None: + pipeline_options.picture_description_options.generation_config[ + "max_new_tokens" + ] = picture_description_max_new_tokens if _should_generate_export_images(image_export_mode, to_formats): pipeline_options.generate_page_images = True pipeline_options.generate_picture_images = True diff --git a/docling/datamodel/chart_extraction_options.py b/docling/datamodel/chart_extraction_options.py index ac7658fcec..067bbb809c 100644 --- a/docling/datamodel/chart_extraction_options.py +++ b/docling/datamodel/chart_extraction_options.py @@ -33,6 +33,7 @@ class ChartExtractionOutputFormat(str, Enum): """ GRANITE_VISION_CHARTS = "granite_vision_charts" + GRANITE_VISION_CHART2CSV = "granite_vision_chart2csv" class ChartExtractionVlmEngineOptions(StagePresetMixin, VlmEngineOptionsMixin): @@ -47,14 +48,13 @@ class ChartExtractionVlmEngineOptions(StagePresetMixin, VlmEngineOptionsMixin): * ``chart2summary`` — generate a natural-language description (default: False) * ``chart2code`` — generate Python code that recreates the chart (default: False) - .. note:: - The ``granite_vision`` V1 preset (ibm-granite/granite-vision-3.3-2b-chart2csv-preview) - was removed in this release. Use ``granite_vision_v4`` (the default) instead. - The last release supporting V1 was 2.x (see the changelog for migration guidance). + The ``granite_vision`` preset uses the older CSV-only Chart2CSV model and + its recommended plain-language prompt. ``granite_vision_v4`` remains the + default and additionally supports summaries and code. Examples:: - # Default preset (granite_vision_v4, Transformers engine) + # Default preset (granite_vision_v4, automatic local engine selection) options = ChartExtractionVlmEngineOptions.from_preset("granite_vision_v4") # Override engine at preset time @@ -127,6 +127,11 @@ def _at_least_one_output(self) -> Self: raise ValueError( "At least one of chart2csv, chart2summary, or chart2code must be True." ) + if self.output_format == ChartExtractionOutputFormat.GRANITE_VISION_CHART2CSV: + if not self.chart2csv or self.chart2summary or self.chart2code: + raise ValueError( + "The granite_vision Chart2CSV preset supports CSV output only." + ) return self def active_prompts(self) -> list[str]: @@ -148,9 +153,15 @@ def active_prompts(self) -> list[str]: from docling.datamodel import stage_model_specs as _stage_model_specs # noqa: E402 +ChartExtractionVlmEngineOptions.register_preset( + _stage_model_specs.CHART_EXTRACTION_GRANITE_VISION +) ChartExtractionVlmEngineOptions.register_preset( _stage_model_specs.CHART_EXTRACTION_GRANITE_VISION_V4 ) +ChartExtractionVlmEngineOptions.register_preset( + _stage_model_specs.CHART_EXTRACTION_GRANITE_VISION_V4_MLX +) # --------------------------------------------------------------------------- @@ -181,10 +192,7 @@ class ChartExtractionModelKind(metaclass=_ChartExtractionModelKindMeta): Use :meth:`ChartExtractionVlmEngineOptions.from_preset` with ``'granite_vision_v4'`` instead. - .. note:: - ``GRANITE_VISION`` (V1) support has been removed. References to - ``ChartExtractionModelKind.GRANITE_VISION`` will resolve to - ``'granite-vision-v4'`` with a deprecation warning. + ``GRANITE_VISION`` selects the older CSV-only Chart2CSV model. """ GRANITE_VISION = "granite-vision" @@ -214,9 +222,9 @@ def __hash__(self) -> int: _members: ClassVar[Dict[str, "_ChartExtractionModelKindMeta"]] = {} - # Map old enum values to new preset IDs (V1 → V4 with deprecation) + # Map legacy enum values to the registered preset IDs. _PRESET_MAP: ClassVar[Dict[str, str]] = { - "granite-vision": "granite_vision_v4", + "granite-vision": "granite_vision", "granite-vision-v4": "granite_vision_v4", } @@ -236,9 +244,8 @@ class ChartExtractionModelOptions(ChartExtractionVlmEngineOptions): For backwards compatibility, instantiating this class emits a ``DeprecationWarning`` and returns a fully functional - ``ChartExtractionVlmEngineOptions`` configured from the ``granite_vision_v4`` - preset. Passing ``model=ChartExtractionModelKind.GRANITE_VISION`` (V1) is - accepted but silently upgraded to V4 with an additional warning. + ``ChartExtractionVlmEngineOptions`` configured from the requested preset. + The default remains ``granite_vision_v4``. """ kind: ClassVar[Literal["chart_extraction"]] = "chart_extraction" # type: ignore[assignment] @@ -268,18 +275,12 @@ def __init__(self, **data: Any) -> None: f"Unknown model {model_str!r}. " f"Valid values: {list(ChartExtractionModelKind._PRESET_MAP)}" ) - if model_str == ChartExtractionModelKind.GRANITE_VISION: - warnings.warn( - "ChartExtractionModelKind.GRANITE_VISION (V1) is no longer supported " - "and has been upgraded to granite_vision_v4.", - DeprecationWarning, - stacklevel=2, - ) - # Bootstrap from the preset so model_spec and engine_options are populated, # then allow the caller's remaining kwargs (chart2csv, etc.) to override. preset_instance = ChartExtractionVlmEngineOptions.from_preset( - "granite_vision_v4" + ChartExtractionModelKind._PRESET_MAP.get( + str(model_val), "granite_vision_v4" + ) ) merged = {**preset_instance.model_dump(), **data} super().__init__(**merged) diff --git a/docling/datamodel/pipeline_options.py b/docling/datamodel/pipeline_options.py index 299bd41625..3f19297e45 100644 --- a/docling/datamodel/pipeline_options.py +++ b/docling/datamodel/pipeline_options.py @@ -1245,7 +1245,7 @@ class CodeFormulaVlmOptions(StagePresetMixin, VlmEngineOptionsMixin, BaseModel): _default_chart_extraction_options = ChartExtractionVlmEngineOptions.from_preset( "granite_vision_v4" ) -"""Default chart extraction options using granite_vision_v4 preset with Transformers runtime.""" +"""Default chart extraction options using granite_vision_v4 with automatic runtime selection.""" # Define an enum for the backend options @@ -1474,7 +1474,7 @@ class ConvertPipelineOptions(PipelineOptions): description=( "Configuration for the chart extraction stage. " "Use ChartExtractionVlmEngineOptions.from_preset('granite_vision_v4') " - "(default) or from_preset('granite_vision') for the V1 model. " + "(default) or from_preset('granite_vision') for the CSV-only model. " "Controls which output formats are generated (chart2csv, chart2summary, chart2code)." ) ), diff --git a/docling/datamodel/service/options.py b/docling/datamodel/service/options.py index 081c155cd8..6c6d332d1f 100644 --- a/docling/datamodel/service/options.py +++ b/docling/datamodel/service/options.py @@ -654,9 +654,14 @@ class ConvertDocumentsOptions(BaseModel): description=( "Preset ID for chart extraction. " 'Use "default" for the admin-controlled default, or a specific preset ' - 'such as "granite_vision_v4" or "granite_vision".' + 'such as "granite_vision_v4", "granite_vision_v4_mlx", or "granite_vision".' ), - examples=["default", "granite_vision_v4", "granite_vision"], + examples=[ + "default", + "granite_vision_v4", + "granite_vision_v4_mlx", + "granite_vision", + ], ), ] = None diff --git a/docling/datamodel/stage_model_specs.py b/docling/datamodel/stage_model_specs.py index 8439bd1142..5803651aa5 100644 --- a/docling/datamodel/stage_model_specs.py +++ b/docling/datamodel/stage_model_specs.py @@ -1959,7 +1959,7 @@ def from_preset( ), scale=2.0, default_engine_type=VlmEngineType.TRANSFORMERS, - stage_options={"output_format": "granite_vision_charts"}, + stage_options={"output_format": "granite_vision_chart2csv"}, ) CHART_EXTRACTION_GRANITE_VISION_V4 = StageModelPreset( @@ -1978,6 +1978,7 @@ def from_preset( trust_remote_code=True, supported_engines={ VlmEngineType.TRANSFORMERS, + VlmEngineType.MLX, VlmEngineType.API_LMSTUDIO, VlmEngineType.API_OLLAMA, VlmEngineType.API_OPENAI, @@ -1989,6 +1990,7 @@ def from_preset( "transformers_model_type": TransformersModelType.AUTOMODEL_IMAGETEXTTOTEXT, }, ), + VlmEngineType.MLX: EngineModelConfig(min_engine_version="0.7.0"), }, api_overrides={ VlmEngineType.API_LMSTUDIO: ApiModelConfig( @@ -2003,6 +2005,16 @@ def from_preset( }, ), scale=2.0, - default_engine_type=VlmEngineType.TRANSFORMERS, + default_engine_type=VlmEngineType.AUTO_INLINE, stage_options={"output_format": "granite_vision_charts"}, ) + +CHART_EXTRACTION_GRANITE_VISION_V4_MLX = StageModelPreset( + preset_id="granite_vision_v4_mlx", + name="Granite-Vision-4.1-4B (MLX)", + description="IBM Granite Vision 4.1-4B chart extraction on Apple Silicon", + model_spec=CHART_EXTRACTION_GRANITE_VISION_V4.model_spec, + scale=CHART_EXTRACTION_GRANITE_VISION_V4.scale, + default_engine_type=VlmEngineType.MLX, + stage_options=CHART_EXTRACTION_GRANITE_VISION_V4.stage_options, +) diff --git a/docling/models/picture_description_base_model.py b/docling/models/picture_description_base_model.py index 033f39aa91..3b002d8c4b 100644 --- a/docling/models/picture_description_base_model.py +++ b/docling/models/picture_description_base_model.py @@ -58,7 +58,11 @@ def __init__( self.images_scale = options.scale def is_processable(self, doc: DoclingDocument, element: NodeItem) -> bool: - return self.enabled and isinstance(element, PictureItem) + return ( + self.enabled + and isinstance(element, PictureItem) + and (element.meta is None or element.meta.description is None) + ) def _annotate_images( self, images: Iterable[Image.Image] diff --git a/docling/models/stages/chart_extraction/granite_vision.py b/docling/models/stages/chart_extraction/granite_vision.py index f6bcf49639..6b86f84a5d 100644 --- a/docling/models/stages/chart_extraction/granite_vision.py +++ b/docling/models/stages/chart_extraction/granite_vision.py @@ -122,7 +122,23 @@ def is_processable(self, doc: DoclingDocument, element: NodeItem) -> bool: ): return False main_pred = element.meta.classification.get_main_prediction() - return main_pred.class_name in SUPPORTED_CHART_TYPES + return main_pred.class_name in SUPPORTED_CHART_TYPES and any( + self._needs_prompt(element, prompt) + for prompt in self.options.active_prompts() + ) + + @staticmethod + def _needs_prompt(item: PictureItem, prompt: str) -> bool: + meta = item.meta + if meta is None: + return True + if prompt == "": + return meta.tabular_chart is None + if prompt == "": + return meta.description is None + if prompt == "": + return meta.code is None + return True def _resolve_runtime_engine_type(self) -> VlmEngineType: selected_engine_type = getattr(self.engine, "selected_engine_type", None) @@ -166,11 +182,22 @@ def __call__( stop_strings = list(model_spec.stop_strings) extra_generation_config = model_spec.get_runtime_input_extra_config(engine_type) - # Build a flat batch: image x prompt, keeping them in sync + # Request only chart fields that the VLM did not already provide. batch_inputs: list[VlmEngineInput] = [] - for image in images: + requests: list[tuple[int, str]] = [] + for img_idx, (item, image) in enumerate(zip(elements, images)): for prompt in active_prompts: - wire_prompt = _NL_PROMPT_MAP.get(prompt, prompt) if use_nl else prompt + if not self._needs_prompt(item, prompt): + continue + if ( + self.options.output_format + == ChartExtractionOutputFormat.GRANITE_VISION_CHART2CSV + ): + wire_prompt = model_spec.prompt + else: + wire_prompt = ( + _NL_PROMPT_MAP.get(prompt, prompt) if use_nl else prompt + ) batch_inputs.append( VlmEngineInput( image=image, @@ -181,43 +208,39 @@ def __call__( extra_generation_config=extra_generation_config, ) ) + requests.append((img_idx, prompt)) + + if not batch_inputs: + yield from elements + return if self.engine is None: raise RuntimeError("Engine not initialized") outputs = list(self.engine.predict_batch(batch_inputs)) - n_prompts = len(active_prompts) - for img_idx, item in enumerate(elements): - if not isinstance(item, PictureItem): - yield item - continue + handler = _OUTPUT_FORMAT_HANDLERS.get(self.options.output_format) + if handler is None: + _log.error( + f"No handler registered for output_format " + f"{self.options.output_format!r}; skipping chart extraction." + ) + yield from elements + return + for (img_idx, prompt), output in zip(requests, outputs): + item = elements[img_idx] if item.meta is None or not isinstance(item.meta, PictureMeta): item.meta = PictureMeta() + _log.debug( + f"chart extraction [{prompt}] image {img_idx}: {output.text[:120]}" + ) + try: + handler(prompt, output.text, item) + except Exception as exc: + _log.error(f"Failed to process [{prompt}] for image {img_idx}: {exc}") - handler = _OUTPUT_FORMAT_HANDLERS.get(self.options.output_format) - if handler is None: - _log.error( - f"No handler registered for output_format " - f"{self.options.output_format!r}; skipping image {img_idx}." - ) - yield item - continue - - for prompt_idx, prompt in enumerate(active_prompts): - result = outputs[img_idx * n_prompts + prompt_idx].text - _log.debug( - f"chart extraction [{prompt}] image {img_idx}: {result[:120]}" - ) - try: - handler(prompt, result, item) - except Exception as exc: - _log.error( - f"Failed to process [{prompt}] for image {img_idx}: {exc}" - ) - - yield item + yield from elements def __del__(self) -> None: if self.engine is not None: @@ -269,6 +292,7 @@ def _handle_granite_vision_charts(prompt: str, result: str, item: PictureItem) - _OUTPUT_FORMAT_HANDLERS: Dict[ChartExtractionOutputFormat, _ChartOutputHandler] = { ChartExtractionOutputFormat.GRANITE_VISION_CHARTS: _handle_granite_vision_charts, + ChartExtractionOutputFormat.GRANITE_VISION_CHART2CSV: _handle_granite_vision_charts, } diff --git a/docling/models/stages/picture_classifier/document_picture_classifier.py b/docling/models/stages/picture_classifier/document_picture_classifier.py index 9e17d3a9b5..5c55825a82 100644 --- a/docling/models/stages/picture_classifier/document_picture_classifier.py +++ b/docling/models/stages/picture_classifier/document_picture_classifier.py @@ -121,7 +121,11 @@ def is_processable(self, doc: DoclingDocument, element: NodeItem) -> bool: bool True if the element is a PictureItem and processing is enabled; False otherwise. """ - return self.enabled and isinstance(element, PictureItem) + return ( + self.enabled + and isinstance(element, PictureItem) + and (element.meta is None or element.meta.classification is None) + ) def __call__( self, diff --git a/docling/pipeline/vlm_pipeline.py b/docling/pipeline/vlm_pipeline.py index ff67a579fa..09f3e39068 100644 --- a/docling/pipeline/vlm_pipeline.py +++ b/docling/pipeline/vlm_pipeline.py @@ -84,9 +84,31 @@ def __init__(self, pipeline_options: VlmPipelineOptions): else: self._initialize_legacy_vlm_models(pipeline_options) - self.enrichment_pipe: list = [ - # Other models working on `NodeItem` elements in the DoclingDocument - ] + def _release_page_resources(self, page: Page) -> None: + scales: set[float] = set() + if self.pipeline_options.do_picture_classification: + scales.add(2.0) + if self.pipeline_options.do_picture_description: + scales.add(self.pipeline_options.picture_description_options.scale) + if self.pipeline_options.do_chart_extraction: + scales.add(2.0) + if scales and page.size is not None: + if page._backend is not None and page._backend.is_valid(): + # Enrichment runs after page assembly and its iterator is closed. + # Cache the scales needed to crop pictures after backend cleanup. + for scale in scales: + page.get_image(scale=scale) + page._backend.unload() + page._backend = None + page.parsed_page = None + return + super()._release_page_resources(page) + + def _unload(self, conv_res: ConversionResult) -> ConversionResult: + result = super()._unload(conv_res) + for page in conv_res.pages: + page._image_cache = {} + return result def _initialize_new_runtime_system( self, pipeline_options: VlmPipelineOptions diff --git a/docling/utils/model_downloader.py b/docling/utils/model_downloader.py index 65dcdda9d7..1bfce4344b 100644 --- a/docling/utils/model_downloader.py +++ b/docling/utils/model_downloader.py @@ -218,10 +218,20 @@ def download_models( ) if with_granite_chart_extraction: - _log.warning( - "with_granite_chart_extraction=True: the Granite Vision V1 chart extraction " - "model (granite-vision-3.3-2b-chart2csv-preview) is no longer supported. " - "Use with_granite_chart_extraction_v4=True instead." + from docling.datamodel.chart_extraction_options import ( + ChartExtractionVlmEngineOptions, + ) + + preset = ChartExtractionVlmEngineOptions.get_preset("granite_vision") + repo_id = preset.model_spec.get_repo_id(preset.default_engine_type) + revision = preset.model_spec.get_revision(preset.default_engine_type) + _log.info(f"Downloading Granite Vision 3.3 Chart2CSV model ({repo_id})...") + download_hf_model( + repo_id=repo_id, + revision=revision, + local_dir=output_dir / repo_id.replace("/", "--"), + force=force, + progress=progress, ) if with_granite_chart_extraction_v4: diff --git a/docs/examples/chart_extraction.py b/docs/examples/chart_extraction.py index 146c23850b..0a0eec739a 100644 --- a/docs/examples/chart_extraction.py +++ b/docs/examples/chart_extraction.py @@ -22,7 +22,8 @@ # Notes # - Setting `do_chart_extraction=True` automatically enables picture classification. # - Supported chart types: bar chart, pie chart, line chart. -# - The default preset uses the local Transformers runtime (granite_vision_v4). +# - The default preset selects MLX on compatible Apple Silicon systems and +# otherwise uses Transformers (granite_vision_v4). # Pass --lmstudio to use the GGUF model served by LM Studio instead. # %% diff --git a/docs/usage/vision_models.md b/docs/usage/vision_models.md index 3e8bd902b3..8759d0a2dd 100644 --- a/docs/usage/vision_models.md +++ b/docs/usage/vision_models.md @@ -43,6 +43,33 @@ For running Docling using local models with the `VlmPipeline`: doc = converter.convert(source="FILE").document ``` +To extract chart data after VLM conversion, enable chart enrichment. Docling +classifies pictures before chart extraction and skips fields already present in +the VLM output. + +```bash +docling --from image --to dclx IMAGE.png --pipeline vlm \ + --vlm-model mineru2_pro --enrich-chart-extraction +``` + +The Python SDK uses `VlmPipelineOptions(do_chart_extraction=True)` for the same +behavior. Picture classification is enabled automatically when chart extraction +is requested. Picture description and classification can also be enabled with +`--enrich-picture-description` and `--enrich-picture-classes`. + +The default chart model is Granite Vision 4.1 4B. On Apple Silicon, it uses +MLX when `mlx-vlm>=0.7.0` and MPS are available; otherwise it uses Transformers. +To use the older CSV-only +`ibm-granite/granite-vision-3.3-2b-chart2csv-preview` model, add +`--chart-extraction-preset granite_vision` to the command. In Python, set +`chart_extraction_options=ChartExtractionVlmEngineOptions.from_preset("granite_vision")` +on `VlmPipelineOptions` or `PdfPipelineOptions`. + +To require MLX explicitly, select `--chart-extraction-preset granite_vision_v4_mlx`. +This preset loads the same official +`ibm-granite/granite-vision-4.1-4b` repository. In Python, use +`ChartExtractionVlmEngineOptions.from_preset("granite_vision_v4_mlx")`. + ## Available local models By default, the vision-language models are running locally. diff --git a/pyproject.toml b/pyproject.toml index 5e07452bb7..ec982c16eb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -269,7 +269,7 @@ models-vlm-inline = [ 'einops>=0.8.1,<1.0.0', 'open-clip-torch>=3.2.0,<4.0.0', 'timm>=1.0.22,<2.0.0', - 'mlx-vlm>=0.6.17,<1.0.0 ; python_version >= "3.10" and sys_platform == "darwin" and platform_machine == "arm64"', + 'mlx-vlm>=0.7.0,<1.0.0 ; python_version >= "3.10" and sys_platform == "darwin" and platform_machine == "arm64"', 'qwen-vl-utils>=0.0.11', 'peft>=0.18.1', ] diff --git a/tests/test_chart_extraction_tabledata.py b/tests/test_chart_extraction_tabledata.py index c9aa0db1c8..745f660d97 100644 --- a/tests/test_chart_extraction_tabledata.py +++ b/tests/test_chart_extraction_tabledata.py @@ -1,17 +1,110 @@ # SPDX-FileCopyrightText: The Docling Contributors # SPDX-License-Identifier: MIT -"""Regression tests for chart CSV table header semantics.""" +"""Regression tests for chart extraction presets and CSV table semantics.""" + +import sys +from types import ModuleType import pandas as pd import pytest +from docling_core.types.doc import ( + DescriptionMetaField, + DoclingDocument, + PictureClassificationMetaField, + PictureMeta, + TableData, + TabularChartMetaField, +) +from docling_core.types.doc.document import PictureClassificationPrediction +from PIL import Image +from docling.datamodel.accelerator_options import AcceleratorOptions +from docling.datamodel.base_models import ItemAndImageEnrichmentElement +from docling.datamodel.chart_extraction_options import ChartExtractionVlmEngineOptions +from docling.datamodel.vlm_engine_options import ( + AutoInlineVlmEngineOptions, + MlxVlmEngineOptions, +) +from docling.models.inference_engines.vlm.auto_inline_engine import ( + AutoInlineVlmEngine, +) +from docling.models.inference_engines.vlm.base import VlmEngineType from docling.models.stages.chart_extraction.granite_vision import ( + ChartExtractionVlmEngineModel, _dataframe_to_tabledata, _extract_csv_to_dataframe, ) +def test_granite_vision_v4_mlx_preset_uses_official_model() -> None: + mlx_options = ChartExtractionVlmEngineOptions.from_preset("granite_vision_v4_mlx") + default_options = ChartExtractionVlmEngineOptions.from_preset("granite_vision_v4") + + assert isinstance(mlx_options.engine_options, MlxVlmEngineOptions) + assert isinstance(default_options.engine_options, AutoInlineVlmEngineOptions) + assert mlx_options.model_spec.is_engine_supported(VlmEngineType.MLX) + assert mlx_options.model_spec.get_engine_config(VlmEngineType.MLX).repo_id == ( + "ibm-granite/granite-vision-4.1-4b" + ) + assert mlx_options.model_spec.get_engine_config(VlmEngineType.MLX).revision == ( + default_options.model_spec.revision + ) + assert mlx_options.output_format == default_options.output_format + assert ( + default_options.model_spec.get_engine_config( + VlmEngineType.MLX + ).min_engine_version + == "0.7.0" + ) + assert "granite_vision_v4_mlx" in ChartExtractionVlmEngineOptions.list_preset_ids() + + +@pytest.mark.parametrize( + ("system", "device", "mlx_version_ok", "expected_engine"), + [ + ("Darwin", "mps", True, VlmEngineType.MLX), + ("Darwin", "mps", False, VlmEngineType.TRANSFORMERS), + ("Darwin", "cpu", True, VlmEngineType.TRANSFORMERS), + ("Linux", "cpu", True, VlmEngineType.TRANSFORMERS), + ], +) +def test_granite_vision_v4_auto_selects_local_engine( + monkeypatch: pytest.MonkeyPatch, + system: str, + device: str, + mlx_version_ok: bool, + expected_engine: VlmEngineType, +) -> None: + options = ChartExtractionVlmEngineOptions.from_preset("granite_vision_v4") + assert isinstance(options.engine_options, AutoInlineVlmEngineOptions) + engine = AutoInlineVlmEngine( + options=options.engine_options, + accelerator_options=AcceleratorOptions(), + artifacts_path=None, + ) + engine.model_spec = options.model_spec + + monkeypatch.setattr("platform.system", lambda: system) + monkeypatch.setattr( + "docling.models.inference_engines.vlm.auto_inline_engine.decide_device", + lambda *args, **kwargs: device, + ) + monkeypatch.setitem(sys.modules, "mlx_vlm", ModuleType("mlx_vlm")) + + def version_satisfied(engine_type: VlmEngineType, min_version: str | None) -> bool: + assert engine_type == VlmEngineType.MLX + assert min_version == "0.7.0" + return mlx_version_ok + + monkeypatch.setattr( + "docling.models.inference_engines.vlm.auto_inline_engine.engine_version_satisfied", + version_satisfied, + ) + + assert engine._select_engine() == expected_engine + + @pytest.mark.parametrize( ("csv_text", "header_count", "row_header_coords"), [ @@ -94,3 +187,107 @@ def test_empty_chart_table_has_no_headers() -> None: assert table.num_rows == 0 assert table.num_cols == 0 assert table.table_cells == [] + + +def test_chart_enrichment_runs_only_missing_outputs_after_classification() -> None: + class Engine: + def __init__(self) -> None: + self.prompts: list[str] = [] + + def predict_batch(self, inputs): + self.prompts = [item.prompt for item in inputs] + return [type("Output", (), {"text": "```python\npass\n```"})()] + + def cleanup(self) -> None: + pass + + options = ChartExtractionVlmEngineOptions.from_preset("granite_vision_v4") + options.chart2summary = True + options.chart2code = True + model = ChartExtractionVlmEngineModel.__new__(ChartExtractionVlmEngineModel) + model.enabled = True + model.options = options + engine = Engine() + model.engine = engine + + chart_data = TabularChartMetaField( + chart_data=TableData(num_rows=0, num_cols=0, table_cells=[]) + ) + description = DescriptionMetaField(text="Provided by VLM") + doc = DoclingDocument(name="chart") + picture = doc.add_picture() + picture.meta = PictureMeta(tabular_chart=chart_data, description=description) + assert not model.is_processable(doc, picture) + + picture.meta.classification = PictureClassificationMetaField( + predictions=[PictureClassificationPrediction(class_name="bar_chart")] + ) + assert model.is_processable(doc, picture) + + image = Image.new("RGB", (10, 10), "white") + result = list( + model( + doc, + [ItemAndImageEnrichmentElement(item=picture, image=image)], + ) + ) + + assert result == [picture] + assert engine.prompts == [""] + assert picture.meta.tabular_chart is chart_data + assert picture.meta.description is description + assert picture.meta.code is not None + + +def test_legacy_chart_preset_uses_its_csv_prompt_and_parser() -> None: + class Engine: + def __init__(self) -> None: + self.prompts: list[str] = [] + + def predict_batch(self, inputs): + self.prompts = [item.prompt for item in inputs] + return [type("Output", (), {"text": "Category,Value\nNorth,10"})()] + + def cleanup(self) -> None: + pass + + options = ChartExtractionVlmEngineOptions.from_preset("granite_vision") + assert ( + options.model_spec.default_repo_id + == "ibm-granite/granite-vision-3.3-2b-chart2csv-preview" + ) + with pytest.raises(ValueError, match="supports CSV output only"): + ChartExtractionVlmEngineOptions.from_preset( + "granite_vision", chart2summary=True + ) + + model = ChartExtractionVlmEngineModel.__new__(ChartExtractionVlmEngineModel) + model.enabled = True + model.options = options + engine = Engine() + model.engine = engine + + doc = DoclingDocument(name="chart") + picture = doc.add_picture() + picture.meta = PictureMeta( + classification=PictureClassificationMetaField( + predictions=[PictureClassificationPrediction(class_name="bar_chart")] + ) + ) + assert model.is_processable(doc, picture) + + result = list( + model( + doc, + [ + ItemAndImageEnrichmentElement( + item=picture, image=Image.new("RGB", (10, 10), "white") + ) + ], + ) + ) + + assert result == [picture] + assert engine.prompts == [options.model_spec.prompt] + assert picture.meta.tabular_chart is not None + assert picture.meta.tabular_chart.chart_data.num_rows == 2 diff --git a/tests/test_cli.py b/tests/test_cli.py index 914179127f..a2117563c4 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -1305,8 +1305,36 @@ def test_parse_page_range_is_shared_with_convert_remote(): _parse_page_range("4-2") +@pytest.mark.parametrize( + ("preset_id", "device", "repo_id", "engine_type"), + [ + ( + "granite_vision", + "cpu", + "ibm-granite/granite-vision-3.3-2b-chart2csv-preview", + "transformers", + ), + ( + "granite_vision_v4_mlx", + "mps", + "ibm-granite/granite-vision-4.1-4b", + "mlx", + ), + ( + "granite_vision_v4", + "mps", + "ibm-granite/granite-vision-4.1-4b", + "auto_inline", + ), + ], +) def test_cli_passes_accelerator_options_to_vlm_pipeline( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + preset_id: str, + device: str, + repo_id: str, + engine_type: str, ) -> None: captured_pipeline_options: VlmPipelineOptions | None = None @@ -1351,18 +1379,42 @@ def convert_all( "--pipeline", "vlm", "--device", - "cpu", + device, "--num-threads", "7", + "--enrich-chart-extraction", + "--chart-extraction-preset", + preset_id, + "--enrich-picture-description", + "--picture-description-max-new-tokens", + "128", ], ) assert result.exit_code == 0 assert captured_pipeline_options is not None - assert captured_pipeline_options.accelerator_options.device == AcceleratorDevice.CPU + assert captured_pipeline_options.accelerator_options.device == AcceleratorDevice( + device + ) assert captured_pipeline_options.accelerator_options.num_threads == 7 assert captured_pipeline_options.generate_page_images is True assert captured_pipeline_options.generate_picture_images is True + assert captured_pipeline_options.do_chart_extraction is True + assert ( + captured_pipeline_options.chart_extraction_options.model_spec.default_repo_id + == repo_id + ) + assert ( + captured_pipeline_options.chart_extraction_options.engine_options.engine_type.value + == engine_type + ) + assert captured_pipeline_options.do_picture_description is True + assert ( + captured_pipeline_options.picture_description_options.generation_config[ + "max_new_tokens" + ] + == 128 + ) def _capture_cli_engine_options(monkeypatch, extra_args, tmp_path, option_name): diff --git a/tests/test_model_downloader_chart.py b/tests/test_model_downloader_chart.py new file mode 100644 index 0000000000..8bf6764d1d --- /dev/null +++ b/tests/test_model_downloader_chart.py @@ -0,0 +1,37 @@ +# SPDX-FileCopyrightText: The Docling Contributors +# SPDX-License-Identifier: MIT + +from pathlib import Path + +import pytest + +from docling.utils import model_downloader + + +def test_legacy_chart_download_uses_registered_checkpoint( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + downloads: list[dict] = [] + monkeypatch.setattr( + model_downloader, + "download_hf_model", + lambda **kwargs: downloads.append(kwargs), + ) + + model_downloader.download_models( + output_dir=tmp_path, + with_layout=False, + with_tableformer=False, + with_code_formula=False, + with_picture_classifier=False, + with_rapidocr=False, + with_granite_chart_extraction=True, + ) + + assert len(downloads) == 1 + assert downloads[0]["repo_id"] == ( + "ibm-granite/granite-vision-3.3-2b-chart2csv-preview" + ) + assert downloads[0]["local_dir"] == ( + tmp_path / "ibm-granite--granite-vision-3.3-2b-chart2csv-preview" + ) diff --git a/tests/test_vlm_enrichment.py b/tests/test_vlm_enrichment.py new file mode 100644 index 0000000000..4f44a9ced3 --- /dev/null +++ b/tests/test_vlm_enrichment.py @@ -0,0 +1,80 @@ +# SPDX-FileCopyrightText: The Docling Contributors +# SPDX-License-Identifier: MIT + +from types import SimpleNamespace + +import pytest +from docling_core.types.doc import ( + DescriptionMetaField, + DoclingDocument, + PictureClassificationMetaField, + PictureMeta, +) +from docling_core.types.doc.document import PictureClassificationPrediction + +from docling.datamodel.pipeline_options import VlmPipelineOptions +from docling.models.picture_description_base_model import PictureDescriptionBaseModel +from docling.models.stages.chart_extraction.granite_vision import ( + ChartExtractionVlmEngineModel, +) +from docling.models.stages.picture_classifier.document_picture_classifier import ( + DocumentPictureClassifier, +) +from docling.pipeline.base_pipeline import ConvertPipeline +from docling.pipeline.vlm_pipeline import VlmPipeline + + +def test_vlm_chart_extraction_runs_after_picture_classification( + monkeypatch: pytest.MonkeyPatch, +) -> None: + description_model = object() + + def init_classifier(self, *, enabled: bool, **kwargs) -> None: + self.enabled = enabled + + def init_chart(self, *, enabled: bool, **kwargs) -> None: + self.enabled = enabled + self.engine = None + + monkeypatch.setattr(DocumentPictureClassifier, "__init__", init_classifier) + monkeypatch.setattr(ChartExtractionVlmEngineModel, "__init__", init_chart) + monkeypatch.setattr( + ConvertPipeline, + "_get_picture_description_model", + lambda self, artifacts_path=None: description_model, + ) + monkeypatch.setattr( + VlmPipeline, + "_initialize_new_runtime_system", + lambda self, pipeline_options: None, + ) + + pipeline = VlmPipeline(VlmPipelineOptions(do_chart_extraction=True)) + + assert isinstance(pipeline.enrichment_pipe[0], DocumentPictureClassifier) + assert pipeline.enrichment_pipe[0].enabled is True + assert pipeline.enrichment_pipe[1] is description_model + assert isinstance(pipeline.enrichment_pipe[2], ChartExtractionVlmEngineModel) + + +def test_picture_enrichments_skip_metadata_already_in_vlm_output() -> None: + doc = DoclingDocument(name="picture") + picture = doc.add_picture() + classifier = DocumentPictureClassifier.__new__(DocumentPictureClassifier) + classifier.enabled = True + description_model = SimpleNamespace(enabled=True) + + assert classifier.is_processable(doc, picture) + assert PictureDescriptionBaseModel.is_processable(description_model, doc, picture) + + picture.meta = PictureMeta( + classification=PictureClassificationMetaField( + predictions=[PictureClassificationPrediction(class_name="bar_chart")] + ), + description=DescriptionMetaField(text="Provided by VLM"), + ) + + assert not classifier.is_processable(doc, picture) + assert not PictureDescriptionBaseModel.is_processable( + description_model, doc, picture + ) diff --git a/tests/test_vlm_pipeline_streaming.py b/tests/test_vlm_pipeline_streaming.py index 3752fbcc2c..674dc5d112 100644 --- a/tests/test_vlm_pipeline_streaming.py +++ b/tests/test_vlm_pipeline_streaming.py @@ -154,6 +154,7 @@ def _run_pipeline( random_access: bool = False, failed_page_nos: set[int] | None = None, document_timeout: float | None = None, + do_chart_extraction: bool = False, ): tracker = _Tracker() backend = ( @@ -168,6 +169,9 @@ def _run_pipeline( generate_page_images=generate_page_images, generate_picture_images=generate_picture_images, images_scale=1.0, + do_picture_classification=False, + do_picture_description=False, + do_chart_extraction=do_chart_extraction, vlm_options=InlineVlmOptions( prompt="", repo_id="test", @@ -218,6 +222,24 @@ def test_vlm_streams_out_of_order_pages_and_releases_each_batch(monkeypatch) -> assert all(page.image is None for page in conv_res.document.pages.values()) +def test_vlm_keeps_chart_crops_available_after_page_backends_close() -> None: + conv_res, tracker, _backend = _run_pipeline( + page_nos=[5], + force_backend_text=False, + generate_page_images=False, + generate_picture_images=False, + tag="picture", + do_chart_extraction=True, + ) + + assert tracker.live == 0 + assert conv_res.document.pictures + picture = conv_res.document.pictures[0] + page = conv_res.pages[0] + assert page._backend is None + assert page.get_image(scale=2.0, cropbox=picture.prov[0].bbox) is not None + + def test_vlm_uses_indexed_loading_for_random_access_backends(monkeypatch) -> None: monkeypatch.setattr(settings.perf, "page_batch_size", 2) @@ -291,6 +313,9 @@ def test_vlm_text_response_keeps_absolute_page_number_after_concatenation() -> N generate_page_images=False, generate_picture_images=False, images_scale=1.0, + do_picture_classification=False, + do_picture_description=False, + do_chart_extraction=False, vlm_options=InlineVlmOptions( prompt="", repo_id="test", @@ -342,6 +367,9 @@ def test_chandra_page_assembly_preserves_source_provenance_and_reports_bad_respo generate_page_images=True, generate_picture_images=False, images_scale=1.0, + do_picture_classification=False, + do_picture_description=False, + do_chart_extraction=False, vlm_options=InlineVlmOptions( prompt="", repo_id="test", diff --git a/uv.lock b/uv.lock index 0baaeeaf1b..85f09e63c7 100644 --- a/uv.lock +++ b/uv.lock @@ -2345,7 +2345,7 @@ requires-dist = [ { name = "lxml", marker = "extra == 'format-xml-jats'", specifier = ">=4.0.0,<7.0.0" }, { name = "mail-parser", marker = "extra == 'format-email'", specifier = ">=4.1.4,<5.0.0" }, { name = "marko", marker = "extra == 'format-markdown'", specifier = ">=2.1.2,<3.0.0" }, - { name = "mlx-vlm", marker = "python_full_version >= '3.10' and platform_machine == 'arm64' and sys_platform == 'darwin' and extra == 'models-vlm-inline'", specifier = ">=0.6.17,<1.0.0" }, + { name = "mlx-vlm", marker = "python_full_version >= '3.10' and platform_machine == 'arm64' and sys_platform == 'darwin' and extra == 'models-vlm-inline'", specifier = ">=0.7.0,<1.0.0" }, { name = "mlx-whisper", marker = "python_full_version >= '3.10' and platform_machine == 'arm64' and sys_platform == 'darwin' and extra == 'format-audio'", specifier = ">=0.4.3" }, { name = "nemotron-ocr", marker = "python_full_version == '3.12.*' and platform_machine == 'x86_64' and sys_platform == 'linux' and extra == 'feat-ocr-nemotron'", specifier = ">=2.0.0" }, { name = "numba", marker = "extra == 'format-audio'", specifier = ">=0.63.0" },