| import json |
| import sys |
| import warnings |
|
|
| import pandas as pd |
| import joblib |
|
|
| try: |
| import cuml |
| except ImportError: |
| pass |
|
|
| try: |
| import xgboost |
| except ImportError: |
| pass |
|
|
| try: |
| import catboost |
| except ImportError: |
| pass |
|
|
| try: |
| import dense_utils |
| except ImportError: |
| pass |
|
|
| warnings.filterwarnings("ignore", module="sklearn") |
|
|
| MODEL_PATH = "model.joblib" |
| _model = None |
| _label_enc = None |
|
|
|
|
| def load_model(path): |
| data = joblib.load(path) |
| return data["model"], data["label_encoder"] |
|
|
|
|
| def _ensure_model(): |
| global _model, _label_enc |
| if _model is None: |
| _model, _label_enc = load_model(MODEL_PATH) |
| return _model, _label_enc |
|
|
|
|
| def _row_to_dict(r): |
| return { |
| "text": r["text"], |
| "files_count": r.get("files_count", 0), |
| "additions": r.get("additions", 0), |
| "deletions": r.get("deletions", 0), |
| "changed_tests": r.get("changed_tests", 0), |
| "changed_docs": r.get("changed_docs", 0), |
| "changed_source": r.get("changed_source", 0), |
| "has_tests": int(r.get("has_tests", False)), |
| "has_docs": int(r.get("has_docs", False)), |
| "extensions": " ".join(r.get("extensions", [])), |
| "directories": " ".join(r.get("directories", [])), |
| } |
|
|
|
|
| def _build_scores(label_enc, probs): |
| return sorted( |
| zip(label_enc.classes_, probs), |
| key=lambda x: x[1], |
| reverse=True, |
| ) |
|
|
|
|
| def predict( |
| text, |
| files_count=0, |
| additions=0, |
| deletions=0, |
| changed_tests=0, |
| changed_docs=0, |
| changed_source=0, |
| has_tests=False, |
| has_docs=False, |
| extensions=None, |
| directories=None, |
| ): |
| row = { |
| "text": text, |
| "files_count": files_count, |
| "additions": additions, |
| "deletions": deletions, |
| "changed_tests": changed_tests, |
| "changed_docs": changed_docs, |
| "changed_source": changed_source, |
| "has_tests": int(has_tests), |
| "has_docs": int(has_docs), |
| "extensions": " ".join(extensions or []), |
| "directories": " ".join(directories or []), |
| } |
| model, label_enc = _ensure_model() |
| df = pd.DataFrame([row]) |
| pred = model.predict(df)[0] |
| probs = model.predict_proba(df)[0] |
| label = label_enc.inverse_transform([pred])[0] |
| scores = _build_scores(label_enc, probs) |
| return label, scores |
|
|
|
|
| def predict_batch(records): |
| model, label_enc = _ensure_model() |
| df = pd.DataFrame([_row_to_dict(r) for r in records]) |
| preds = model.predict(df) |
| probs = model.predict_proba(df) |
|
|
| results = [] |
| for i in range(len(records)): |
| label = label_enc.inverse_transform([preds[i]])[0] |
| scores = _build_scores(label_enc, probs[i]) |
| results.append({"label": label, "probs": dict(scores)}) |
| return results |
|
|
|
|
| def main(): |
| if len(sys.argv) < 2: |
| print("Usage: python predict.py <message> [files_count] [additions] [deletions]") |
| print(" or: echo '<json>' | python predict.py --stdin") |
| sys.exit(1) |
|
|
| if sys.argv[1] == "--stdin": |
| records = [json.loads(line.strip()) for line in sys.stdin if line.strip()] |
| if records: |
| results = predict_batch(records) |
| for r in results: |
| print(json.dumps(r, ensure_ascii=False)) |
| return |
|
|
| text = sys.argv[1] |
| files_count = int(sys.argv[2]) if len(sys.argv) > 2 else 0 |
| additions = int(sys.argv[3]) if len(sys.argv) > 3 else 0 |
| deletions = int(sys.argv[4]) if len(sys.argv) > 4 else 0 |
| label, scores = predict(text, files_count, additions, deletions) |
|
|
| print(f"Prediction: {label}") |
| print("Top-3:") |
| for cls, prob in scores[:3]: |
| print(f" {cls:>10}: {prob:.1%}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|