"""
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)