git-commits-sorter / predict.py
akaruineko's picture
Upload folder using huggingface_hub
9655fa2 verified
Raw
History Blame Contribute Delete
3.71 kB
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()