Spaces:
Sleeping
Sleeping
Download src/predict.py from kshitiz14/product_attribute: direct link, hf CLI and curl.
- Browser
- Download file 2.17 kB
-
https://huggingface.co/spaces/kshitiz14/product_attribute/resolve/main/src/predict.py
- Command line
-
hf download hf://spaces/kshitiz14/product_attribute/src/predict.py
-
curl -L -o predict.py https://huggingface.co/spaces/kshitiz14/product_attribute/resolve/main/src/predict.py
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() | |