"""Persian NER enrichment for invoice field extraction.""" from __future__ import annotations import logging import re from functools import lru_cache logger = logging.getLogger(__name__) _ner_pipeline = None @lru_cache(maxsize=1) def _get_ner(): global _ner_pipeline if _ner_pipeline is not None: return _ner_pipeline try: from transformers import pipeline logger.info("Loading Persian NER model (HooshvareLab/bert-fa-zwnj-base-ner)...") _ner_pipeline = pipeline( "ner", model="HooshvareLab/bert-fa-zwnj-base-ner", aggregation_strategy="simple", device=-1, ) return _ner_pipeline except Exception as exc: logger.warning("NER model unavailable: %s", exc) return None def extract_organizations(text: str) -> list[str]: ner = _get_ner() if not ner: return [] try: entities = ner(text[:512]) return [e["word"].replace("##", "") for e in entities if e.get("entity_group") in ("B-ORG", "I-ORG", "ORG")] except Exception: return [] def extract_persons(text: str) -> list[str]: ner = _get_ner() if not ner: return [] try: entities = ner(text[:512]) return [e["word"].replace("##", "") for e in entities if e.get("entity_group") in ("B-PER", "I-PER", "PER")] except Exception: return [] def find_tax_ids(text: str) -> list[str]: normalized = text.translate(str.maketrans("۰۱۲۳۴۵۶۷۸۹", "0123456789")) return re.findall(r"(?