alirezaaminzadeh's picture
Upload folder using huggingface_hub
af24ae8 verified
Raw
History Blame Contribute Delete
1.67 kB
"""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"(?<!\d)(\d{10,14})(?!\d)", normalized)