product_attribute / src /predict.py
kshitiz14's picture
Fix Dockerfile for HF Spaces (port 7860), stop tracking model artifact
896a559
Raw History Blame Contribute Delete
2.17 kB
"""
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()