Source code for flow_preprocessing.preprocessing_logic.preprocess

"""
Flow Preprocessor - Main Preprocessing Module

This module provides preprocessing functionality for PageXML datasets with:
- Dependency Injection pattern
- Factory Pattern for converter creation
- Configuration object pattern
- Support for ZIP files and HuggingFace datasets
- Optional GPU-accelerated image segmentation
"""

import gc
from abc import ABC, abstractmethod
from typing import Any, Literal

import datasets
from flow_segmenter import SegmenterBaseConfig, SegmenterConfig
from loguru import logger
from pagexml_hf import XmlConverter
from pydantic import SecretStr, ValidationError

from flow_preprocessing.preprocessing_logic.config import (
    ExportMode,
    PreprocessorConfig,
    ProcessorState,
)
from flow_preprocessing.preprocessing_logic.converter_factory import ConverterFactory
from flow_preprocessing.utils.url_validator import validate_url

# ===============================================================================
# BASE PREPROCESSOR
# ===============================================================================


class Preprocessor(ABC):
    """
    Base preprocessor class with improved OOP design.

    Features:
    - Dependency Injection for better testability
    - Configuration object pattern for cleaner initialization
    - Factory pattern for converter creation
    - Properties for encapsulation
    """

    def __init__(
        self, config: PreprocessorConfig, converter_factory: ConverterFactory | None = None
    ) -> None:
        """
        Initialize the preprocessor.

        :param config: Configuration object containing all settings.
        :param converter_factory: Optional factory for creating converters (for testing/DI).
        """
        self._config = config
        self._converter_factory = converter_factory or ConverterFactory()
        self._state = ProcessorState.INITIALIZED
        self._dataset: datasets.Dataset | None = None
        self._converter: XmlConverter | None = None

        # Initialize segmenter config
        self._segmenter_config = self._initialize_segmenter_config(config.segmenter_config)
        logger.info(f"Preprocessor initialized with config: {config}")
        logger.debug(f"Preprocessor initialized with converter: {self._converter}")
        logger.debug(
            f"Preprocessor initialized with segmentation models: {self._segmenter_config}"
        )

    # ==================== Properties ====================

    @property
    def state(self) -> ProcessorState:
        """Get current processor state."""
        return self._state

    @property
    def config(self) -> PreprocessorConfig:
        """Get configuration."""
        return self._config

    @property
    def dataset(self) -> datasets.Dataset | None:
        """Get the current dataset."""
        return self._dataset

    @property
    def converter(self) -> XmlConverter:
        """
        Get or create the XmlConverter (lazy initialization).

        :return: XmlConverter instance.
        """
        if self._converter is None:
            self._converter = self.create_xmlconverter()
        return self._converter

    # ==================== Abstract Methods ====================

    @abstractmethod
    def create_xmlconverter(self) -> XmlConverter:
        """
        Create an XmlConverter for the specific data source.

        Subclasses must implement this to define their data source.

        :return: An instance of XmlConverter.
        """

    # ==================== Public Methods ====================

    def preprocess(self) -> str:
        """
        Perform preprocessing steps: segmentation (optional) and dataset conversion/upload.

        :return: URL of the uploaded dataset repository.
        :raises Exception: If preprocessing fails.
        """
        try:
            self._set_state(ProcessorState.IN_PROGRESS)

            # Step 1: Segmentation (if enabled) - non-blocking
            if self._config.segment is not None:
                logger.info("Segmentation enabled - running segment_images()...")
                self.segment_images()
                logger.info("Segmentation completed.")

            # Step 2: Convert and upload - non-blocking
            repo_url = self._convert_and_upload()

            self._set_state(ProcessorState.COMPLETED)
            logger.info(f"Success! Dataset available at: {repo_url}")
            return repo_url

        except Exception as e:
            self._set_state(ProcessorState.FAILED)
            logger.error(f"Preprocessing failed: {e}")
            raise

    def segment_images(self) -> None:
        """
        Segment images in the dataset using YOLO or kraken.

        :raises ValueError: If segmenter_config is not provided.
        """
        try:
            self._segment_images()
        except Exception as e:
            logger.error(f"Segmentation failed: {e}")
            self._set_state(ProcessorState.FAILED)
            raise

    # ==================== Private Sync Methods ====================

    def _initialize_segmenter_config(
        self, config: SegmenterConfig | SegmenterBaseConfig | dict | None
    ) -> SegmenterConfig | SegmenterBaseConfig | None:
        """
        Initialize segmenter configuration.

        :param config: Segmenter config as object or dict.
        :return: SegmenterConfig instance or None.
        :raises ValidationError: If config dict is invalid.
        """
        if config is None:
            return None

        try:
            if isinstance(config, dict) and self._config.segment == "yolo":
                return SegmenterConfig(**config)
            elif isinstance(config, dict) and self._config.segment == "kraken":
                return SegmenterBaseConfig(**config)
            elif isinstance(config, (SegmenterBaseConfig, SegmenterConfig)):
                return config
            else:
                return None
        except ValidationError as e:
            logger.error(f"Invalid segmenter_config: {e}")
            self._set_state(ProcessorState.FAILED)
            raise

    def _segment_images(self) -> None:
        """
        Synchronous implementation of image segmentation.

        :raises ValueError: If segmenter_config is not provided.
        """
        segmenter = None

        if self._segmenter_config is None:
            error_msg = "segmenter_config must be provided when segment is True."
            logger.error(f"Preprocessor._segment_images_sync(): {error_msg}")
            self._set_state(ProcessorState.FAILED)
            raise ValueError(error_msg)

        logger.info("Running segmentation...")

        if self._config.segment == "yolo":
            from flow_segmenter import SegmenterYolo

            logger.info("Using YOLO for segmentation.")
            if isinstance(self._segmenter_config, SegmenterBaseConfig):
                self._segmenter_config = SegmenterConfig(**self._segmenter_config.model_dump())
            # Create segmenter (GPU-accelerated if available)
            segmenter = SegmenterYolo(config=self._segmenter_config)
        elif self._config.segment == "kraken":
            from flow_segmenter import SegmenterKrakenLinemasks

            logger.info("Using Kraken for segmentation.")
            segmenter = SegmenterKrakenLinemasks(config=self._segmenter_config)

        # Convert to raw XML for segmentation
        segmented_dataset = self.converter.convert(
            export_mode=ExportMode.RAW_XML.value,
            # default: split_train=None,
            allow_empty=self._config.allow_empty_lines,
            batch_size=self._config.batch_size,
        )

        # Segment the dataset
        if segmenter is not None:
            self._dataset = segmenter.segment_dataset(
                segmented_dataset,
                new_column_name="xml_content",
            )
            logger.debug(
                f"Segmentation completed. Dataset size: {self._dataset.column_names if self._dataset else 'N/A'}"
            )
            logger.debug(f"Segmented dataset: {segmented_dataset.column_names}")
            del segmenter
            gc.collect()
        else:
            logger.error("No valid segmenter found for segmentation.")
            self._set_state(ProcessorState.FAILED)
            raise ValueError("Invalid segmentation method specified.")

        # Reset converter to use new dataset
        self._converter = None

        logger.info("Segmentation completed.")

    def _convert_and_upload(self) -> str:
        """
        Synchronous implementation of convert and upload.

        :return: URL of the uploaded dataset repository.
        """
        logger.info(f"Converting and uploading with export_mode={self._config.export_mode}")
        if self._config.huggingface_token:
            huggingface_token = (
                self._config.huggingface_token.get_secret_value()
                if isinstance(self._config.huggingface_token, SecretStr)
                else self._config.huggingface_token
            )
        else:
            huggingface_token = ""
        logger.debug(f"HuggingFace token provided: {bool(huggingface_token)}")

        result = self.converter.convert_and_upload(
            repo_id=self._config.huggingface_target_repo_name,
            export_mode=(
                self._config.export_mode
                if self._config.export_mode
                else ExportMode.RAW_XML.value
            ),
            token=str(huggingface_token),
            private=(
                self._config.huggingface_target_repo_private
                if self._config.huggingface_target_repo_private
                else False
            ),
            split_train=self._config.split_train_ratio,
            split_seed=self._config.split_seed,
            split_shuffle=self._config.split_shuffle,
            mask_crop=self._config.crop,
            min_width=self._config.min_width_line,
            min_height=self._config.min_height_line,
            allow_empty=self._config.allow_empty_lines,
            batch_size=self._config.batch_size,
            append=self._config.append,
            line_augment=self._config.augmentation_loops,
        )

        return result if result is not None else self._config.huggingface_target_repo_name

    def _set_state(self, state: ProcessorState) -> None:
        """
        Set the processor state.

        :param state: New state to set.
        """
        self._state = state
        logger.debug(f"Preprocessor state changed to: {state.value}")


# ===============================================================================
# CONCRETE IMPLEMENTATIONS
# ===============================================================================


[docs] class ZipPreprocessor(Preprocessor): """ Preprocessor for ZIP files (local or remote). Supports: - Local ZIP files - Remote ZIP files (HTTP/HTTPS URLs) - Automatic source type detection """
[docs] def __init__( self, input_path: str, config: PreprocessorConfig, converter_factory: ConverterFactory | None = None, ) -> None: """ Initialize ZIP preprocessor. :param input_path: Path or URL to ZIP file. :param config: Preprocessor configuration. :param converter_factory: Optional converter factory (for DI/testing). :raises ValueError: If input_path is invalid or a URL fails validation. """ # Validate input if not input_path: raise ValueError("input_path cannot be empty") if not isinstance(input_path, str): raise TypeError("input_path must be a string") # Validate URL if it's a remote URL if input_path.startswith("http://") or input_path.startswith("https://"): validate_url(input_path) self._input_path = input_path super().__init__(config, converter_factory)
def create_xmlconverter(self) -> XmlConverter: """ Create XmlConverter for ZIP source using factory. :return: Configured XmlConverter instance. """ logger.info(f"Creating XmlConverter for ZIP: {self._input_path}") try: converter = self._converter_factory.create_zip_converter( zip_path=self._input_path, parse_xml=self._config.requires_xml_parsing, dataset=self._dataset, ) logger.info("XmlConverter created successfully.") return converter except Exception as e: logger.error(f"Failed to create XmlConverter: {e}") self._set_state(ProcessorState.FAILED) raise ValueError(f"Failed to create XmlConverter: {e}") from e def preprocess(self) -> str: """ Preprocess ZIP file. :return: URL of uploaded dataset. """ logger.info(f"Preprocessing ZIP: {self._input_path}") repo_url = super().preprocess() logger.info(f"Preprocessing of {self._input_path} completed.") return repo_url
[docs] class HuggingFacePreprocessor(Preprocessor): """ Preprocessor for HuggingFace datasets. Supports processing existing HuggingFace datasets and uploading the processed version to a new or existing repository. """
[docs] def __init__( self, input_path: str, config: PreprocessorConfig, converter_factory: ConverterFactory | None = None, ) -> None: """ Initialize HuggingFace preprocessor. :param input_path: HuggingFace repository ID (e.g., 'username/dataset'). :param config: Preprocessor configuration. :param converter_factory: Optional converter factory (for DI/testing). """ self._input_path = input_path super().__init__(config, converter_factory)
def create_xmlconverter(self) -> XmlConverter: """ Create XmlConverter for HuggingFace source using factory. :return: Configured XmlConverter instance. """ logger.info(f"Creating XmlConverter for HuggingFace: {self._input_path}") try: converter = self._converter_factory.create_huggingface_converter( repo_id=self._input_path, token=self._config.huggingface_token, parse_xml=self._config.requires_xml_parsing, dataset=self._dataset, ) logger.info("XmlConverter created successfully.") return converter except Exception as e: logger.error(f"Failed to create XmlConverter: {e}") self._set_state(ProcessorState.FAILED) raise ValueError(f"Failed to create XmlConverter: {e}") from e def preprocess(self) -> str: """ Preprocess HuggingFace dataset. :return: URL of uploaded dataset. """ logger.info(f"Preprocessing HuggingFace dataset: {self._input_path}") repo_url = super().preprocess() logger.info(f"Preprocessing of {self._input_path} completed.") return repo_url
# =============================================================================== # BUILDER PATTERN (Optional - for easier usage) # ===============================================================================
[docs] class PreprocessorBuilder: """ Builder for creating Preprocessor instances with fluent API. Makes it easier to create preprocessors with complex configurations without needing to manually create PreprocessorConfig objects. Example: preprocessor = (PreprocessorBuilder("username/output-dataset") .with_token("hf_token") .with_export_mode("line") .with_segmentation(segmenter_config) .with_split(0.8) .build_for_zip("data.zip")) """
[docs] def __init__(self, huggingface_target_repo_name: str): """ Initialize builder. :param huggingface_target_repo_name: Target HuggingFace repository. """ self._config_dict: dict[str, Any] = { "huggingface_target_repo_name": huggingface_target_repo_name }
def with_token(self, token: str) -> "PreprocessorBuilder": """ Set HuggingFace token. :param token: HuggingFace access token. :return: Builder instance for chaining. """ self._config_dict["huggingface_token"] = token return self def with_export_mode(self, mode: str) -> "PreprocessorBuilder": """ Set export mode. :param mode: Export mode ('line', 'region', 'text', 'window', 'raw_xml'). :return: Builder instance for chaining. """ self._config_dict["export_mode"] = mode return self def with_crop(self, crop: bool = True) -> "PreprocessorBuilder": """ Enable image cropping. :param crop: Whether to crop images. :return: Builder instance for chaining. """ self._config_dict["crop"] = crop return self def with_segmentation( self, segmenter_config: SegmenterConfig | SegmenterBaseConfig | dict, backend: Literal["yolo", "kraken"] | None = None, ) -> "PreprocessorBuilder": """ Enable image segmentation. :param segmenter_config: Segmenter configuration object or dict. :param backend: Segmentation backend ('yolo' or 'kraken') when config is a dict. :return: Builder instance for chaining. """ if backend is None: if isinstance(segmenter_config, SegmenterConfig): backend = "yolo" elif isinstance(segmenter_config, SegmenterBaseConfig): backend = "kraken" else: raise ValueError("backend must be provided when segmenter_config is a dict") self._config_dict["segment"] = backend self._config_dict["segmenter_config"] = segmenter_config return self def with_split( self, ratio: float, seed: int = 42, shuffle: bool = True ) -> "PreprocessorBuilder": """ Configure dataset splitting. :param ratio: Training data ratio (0.0-1.0). :param seed: Random seed for reproducibility. :param shuffle: Whether to shuffle before splitting. :return: Builder instance for chaining. """ self._config_dict["split_train_ratio"] = ratio self._config_dict["split_seed"] = seed self._config_dict["split_shuffle"] = shuffle return self def with_line_filtering( self, min_width: int | None = None, min_height: int | None = None ) -> "PreprocessorBuilder": """ Configure line filtering. :param min_width: Minimum line width in pixels. :param min_height: Minimum line height in pixels. :return: Builder instance for chaining. """ if min_width is not None: self._config_dict["min_width_line"] = min_width if min_height is not None: self._config_dict["min_height_line"] = min_height return self def with_batch_size(self, batch_size: int) -> "PreprocessorBuilder": """ Set batch size for processing. :param batch_size: Batch size for dataset operations. :return: Builder instance for chaining. """ self._config_dict["batch_size"] = batch_size return self def private(self, is_private: bool = True) -> "PreprocessorBuilder": """ Set repository privacy. :param is_private: Whether the output repository should be private. :return: Builder instance for chaining. """ self._config_dict["huggingface_target_repo_private"] = is_private return self def append(self, should_append: bool = True) -> "PreprocessorBuilder": """ Set append mode. :param should_append: Whether to append to existing dataset. :return: Builder instance for chaining. """ self._config_dict["append"] = should_append return self def build_for_zip(self, zip_path: str) -> ZipPreprocessor: """ Build a ZipPreprocessor. :param zip_path: Path or URL to ZIP file. :return: Configured ZipPreprocessor instance. """ config = PreprocessorConfig(**self._config_dict) return ZipPreprocessor(zip_path, config) def build_for_huggingface(self, repo_id: str) -> HuggingFacePreprocessor: """ Build a HuggingFacePreprocessor. :param repo_id: HuggingFace repository ID. :return: Configured HuggingFacePreprocessor instance. """ config = PreprocessorConfig(**self._config_dict) return HuggingFacePreprocessor(repo_id, config)