File size: 2,165 Bytes
896a559
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
61
62
63
64
65
66
67
68
"""
Inference pipeline: ensembles the trained ML classifiers with the
deterministic rules extractor.

Ensemble strategy (per attribute type, per label):
  - If the rules extractor fires on a label (exact lexicon phrase present),
    keep it -- rules are precision-near-1.0 by construction.
  - Otherwise, trust the ML model's prediction for labels the rules missed
    (e.g. slight paraphrases the lexicon doesn't cover), using probability
    threshold 0.5.
This union favors recall on top of the rules' precision, which is the
practical sweet spot for a 61-row training set where the ML model alone
would still be data-starved for some of the rarer labels.
"""
import joblib
from rules_extractor import extract_attributes_rules

_ARTIFACT = None


def _load(model_path="model.joblib"):
    global _ARTIFACT
    if _ARTIFACT is None:
        _ARTIFACT = joblib.load(model_path)
    return _ARTIFACT


def predict_ml(text: str, model_path="model.joblib") -> dict:
    artifact = _load(model_path)
    X = artifact["vectorizer"].transform([text])
    result = {}
    for attr in artifact["attr_types"]:
        clf = artifact["models"][attr]
        mlb = artifact["binarizers"][attr]
        if clf is None:
            result[attr] = []
            continue
        y_pred = clf.predict(X)
        labels = mlb.inverse_transform(y_pred)[0]
        result[attr] = list(labels)
    return result


def predict_ensemble(text: str, model_path="model.joblib") -> dict:
    rule_preds = extract_attributes_rules(text)
    ml_preds = predict_ml(text, model_path)

    final = {}
    for attr in rule_preds:
        merged = list(rule_preds[attr])
        for label in ml_preds.get(attr, []):
            if label not in merged:
                merged.append(label)
        final[attr] = merged
    return final


if __name__ == "__main__":
    import json
    samples = [
        "Sparkly sequin fitted prom gown featuring a deep illusion neckline and open back",
        "Off shoulder satin ball gown with corset bodice and sweep train in royal navy",
    ]
    for s in samples:
        print(s)
        print(json.dumps(predict_ensemble(s), indent=2))
        print()