ChihHsuan-Yang's picture
Add pretty_name, quiet the version-warning wall in the example, and record the root widget copy's checksum
294c751 verified
Raw History Blame Contribute Delete
6.16 kB
"""Load the released metadata-only router and route a few problems.
Runs offline. No model endpoint, no download beyond this repository.
pip install scikit-learn pandas scipy joblib
python predict_example.py
The router sees problem METADATA only -- difficulty, difficulty tier, source,
and domain. It never sees gold answers, correctness, protocol outcomes, or the
oracle label. That boundary is the point of the paper, so this example does not
give you a way to pass those in.
"""
from __future__ import annotations
import json
from pathlib import Path
import warnings
import joblib
import pandas as pd
from scipy import sparse
HERE = Path(__file__).resolve().parent
# The estimator stores integer classes [0..4] and does NOT carry the label names.
# Decode with this file, not with sorted(labels): the order is the paper's fixed
# cost order, and alphabetical decoding disagrees with the released predictions
# on 71 of 423 test rows while looking perfectly plausible.
#: The scikit-learn version this checkpoint was fitted under, recorded in
#: model_metadata.json. Used only to phrase the version note accurately.
FITTED_WITH = "1.8.0"
MAPPING = json.loads((HERE / "label_mapping.json").read_text())
INDEX_TO_LABEL = {int(k): v for k, v in MAPPING["index_to_label"].items()}
def _quiet_version_warning():
"""Silence the version-mismatch warning, after saying it out loud once.
The estimator was pickled under scikit-learn 1.8.0. Loading it under any
other version makes scikit-learn emit an InconsistentVersionWarning for
EVERY unpickled object -- here that is twelve lines of traceback-looking
text, including the reader's own filesystem paths, before a single line of
output. That reads like a broken artifact.
It is not suppressed silently: one plain sentence replaces the wall, so the
reader still learns the fact the warning was trying to convey.
"""
try:
from sklearn.exceptions import InconsistentVersionWarning
except ImportError: # very old scikit-learn: no such class
return
import sklearn
if sklearn.__version__ != FITTED_WITH:
print(
f"note: this checkpoint was fitted with scikit-learn {FITTED_WITH}; "
f"you are running {sklearn.__version__}.\n"
" predict() is expected to work. If you need exact "
"probabilities, match the fitted version or retrain.\n"
)
warnings.filterwarnings("ignore", category=InconsistentVersionWarning)
def load():
return (
joblib.load(HERE / "model.joblib"),
joblib.load(HERE / "feature_builders.joblib"),
)
def build_features(df: pd.DataFrame, builders: dict):
"""Reproduce the training-time feature pipeline exactly.
Column order matters: source one-hot, then domain multi-hot, then the two
scaled numeric columns. The estimator was fitted on 220 features in that
order and will silently mispredict if they are assembled differently.
"""
x_source = builders["source_encoder"].transform(df[["source"]])
# Domains unseen at training time are dropped rather than erroring, which is
# what the training code does; an unknown domain simply contributes nothing.
known = set(builders["domain_binarizer"].classes_.tolist())
x_domain = builders["domain_binarizer"].transform(
[[d for d in row if d in known] for row in df["domain_list"]]
)
x_numeric = builders["scaler"].transform(
df[["difficulty", "difficulty_tier"]].to_numpy(dtype=float)
)
if not sparse.issparse(x_numeric):
x_numeric = sparse.csr_matrix(x_numeric)
return sparse.hstack([x_source, x_domain, x_numeric], format="csr")
def main() -> None:
_quiet_version_warning()
model, builders = load()
# Three made-up problems, in the schema the router expects.
problems = pd.DataFrame(
[
{"problem_id": "demo-easy", "source": "cayley", "difficulty": 1.5,
"difficulty_tier": 1,
"domain_list": ["Mathematics -> Algebra -> Prealgebra -> Simple Equations"]},
{"problem_id": "demo-mid", "source": "HMMT_2", "difficulty": 4.5,
"difficulty_tier": 5,
"domain_list": ["Mathematics -> Discrete Mathematics -> Combinatorics"]},
{"problem_id": "demo-hard", "source": "imo_shortlist", "difficulty": 8.0,
"difficulty_tier": 8,
"domain_list": ["Mathematics -> Number Theory -> Congruences"]},
]
)
X = build_features(problems, builders)
assert X.shape[1] == model.n_features_in_, (
f"built {X.shape[1]} features but the model expects {model.n_features_in_}"
)
indices = model.predict(X)
# predict_proba reaches into attributes that moved between scikit-learn
# versions, so it can raise on an estimator pickled by a different version
# even though predict() works. Fall back to the decision function, which is
# stable, rather than failing the example.
try:
scores = model.predict_proba(X)
score_label = "confidence"
except Exception:
import numpy as np
margins = model.decision_function(X)
e = np.exp(margins - margins.max(axis=1, keepdims=True))
scores = e / e.sum(axis=1, keepdims=True)
score_label = "softmax(margin)"
print("note: predict_proba is unavailable under this scikit-learn "
"version; showing a softmax over the decision function instead.\n")
print(f"{'problem':12s} {'action':14s} {score_label:16s} runner-up")
for pid, idx, row in zip(problems["problem_id"], indices, scores):
order = row.argsort()[::-1]
top, second = int(order[0]), int(order[1])
print(
f"{pid:12s} {INDEX_TO_LABEL[int(idx)]:14s} {row[top]:.3f}"
f" {INDEX_TO_LABEL[second]} ({row[second]:.3f})"
)
print(
"\n'none' means: do not escalate -- no protocol was expected to succeed,"
"\nso spending on collaboration would be wasted. It is a router action,"
"\nnot a fifth protocol."
)
if __name__ == "__main__":
main()