File size: 6,161 Bytes
0181745
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
294c751
 
0181745
 
 
 
 
 
 
 
 
 
294c751
 
 
 
0181745
 
 
 
294c751
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0181745
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
294c751
0181745
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
"""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()