File size: 1,665 Bytes
af24ae8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
"""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)