"""Load the six family classifiers and predict the family of a reasoning trace.
IMPORTANT: the vectorizers were pickled holding a reference to
`preprocessing.simple_tokenizer`, so `preprocessing.py` (shipped in this repo)
must be importable or `joblib.load` raises
`ModuleNotFoundError: No module named 'preprocessing'`.
python load_classifiers.py # runs the built-in demo
"""
import os
import sys
import joblib
import numpy as np
FAMILIES = ["openai", "qwen", "deepseek", "olmo", "glm", "exaone"]
SUBDIR = "tf_idf"
def load(repo_dir: str = "."):
"""Return (classifiers, vectorizers, preprocess_text) keyed by family."""
if repo_dir not in sys.path:
sys.path.insert(0, repo_dir) # so `preprocessing` resolves
from preprocessing import preprocess_text
base = os.path.join(repo_dir, SUBDIR)
clfs = {f: joblib.load(os.path.join(base, f"classifier_{f}.joblib"))
for f in FAMILIES}
vecs = {f: joblib.load(os.path.join(base, f"vectorizer_{f}.joblib"))
for f in FAMILIES}
return clfs, vecs, preprocess_text
def predict(texts, clfs, vecs, preprocess_text):
"""Family probabilities and argmax prediction for a list of raw traces.
Each classifier is an independent one-vs-rest model, so the six probabilities
do NOT sum to 1. `argmax` is a reasonable combination rule but not a
calibrated posterior.
"""
pre = [preprocess_text(t) for t in texts]
probs = np.zeros((len(pre), len(FAMILIES)))
for j, fam in enumerate(FAMILIES):
probs[:, j] = clfs[fam].predict_proba(vecs[fam].transform(pre))[:, 1]
pred = np.asarray(FAMILIES)[probs.argmax(axis=1)]
return pred, probs
def download(repo_id: str = "SupritiVijay/classifiers_model_provenance") -> str:
"""Grab the repo locally and return the path."""
from huggingface_hub import snapshot_download
return snapshot_download(repo_id)
def main() -> int:
repo_dir = sys.argv[1] if len(sys.argv) > 1 else os.path.dirname(
os.path.abspath(__file__))
print(f"loading from {repo_dir}")
clfs, vecs, preprocess_text = load(repo_dir)
demo = [
"\nWe need to compute the expected value. Let me define dp[i] as "
"the answer for the first i elements. Then dp[i] = dp[i-1] + a[i].\n"
"\nHere is the solution:\n```python\ndef solve(a):\n"
" return sum(a)\n```",
"\nOkay, so the user wants a function that reverses a string. "
"That's straightforward. I'll use slicing.\n\n"
"```python\ndef rev(s):\n return s[::-1]\n```",
]
pred, probs = predict(demo, clfs, vecs, preprocess_text)
print(f"\n{'':4s} " + " ".join(f"{f:>9s}" for f in FAMILIES) + " -> pred")
for i in range(len(demo)):
row = " ".join(f"{p:9.3f}" for p in probs[i])
print(f" {i} {row} -> {pred[i]}")
print("\n(one-vs-rest probabilities; they do not sum to 1)")
return 0
if __name__ == "__main__":
sys.exit(main())