|
| 1 | +# SPDX-FileCopyrightText: Copyright contributors to the cruq.ai project |
| 2 | +# SPDX-License-Identifier: Apache-2.0 |
| 3 | + |
| 4 | +""" |
| 5 | +cruq self-consistency cascade router. |
| 6 | +
|
| 7 | +The router sees only the prompt string. It probes the cheapest capable model |
| 8 | +(qwen3-235b-a22b) K times at temperature>0 and measures self-consistency: the |
| 9 | +fraction of the K samples that agree on the final \\boxed{} answer. High |
| 10 | +agreement is a strong, prompt-independent confidence signal, so: |
| 11 | +
|
| 12 | + - consistency >= tau -> keep qwen's majority-vote answer (cost = K cheap probes) |
| 13 | + - consistency < tau -> escalate to deepseek-v4-flash (a stronger model) |
| 14 | +
|
| 15 | +This is an *inference-time* signal, not a prompt classifier. Across the whole |
| 16 | +research programme, no prompt-only selector (lexical difficulty, per-model |
| 17 | +P(correct) heads, calibrated thresholds, or a domain classifier) beat the best |
| 18 | +single model, because per-query difficulty is close to unpredictable from prompt |
| 19 | +text. Model-agreement at inference is what breaks that ceiling. |
| 20 | +
|
| 21 | +Compliance: nothing here is fit on RouterArena data. K and tau are priors set in |
| 22 | +the config; the probe model and escalation model are fixed choices. No labels, |
| 23 | +metadata, or ground truth are read at decision time. |
| 24 | +
|
| 25 | +Reproducibility: the K qwen probes per query are cached to |
| 26 | +phase2/data/qwen_sc_full.jsonl by phase2/sample_qwen_sc_full.py. This class reads |
| 27 | +that cache to make the same keep/escalate decision the submitted prediction file |
| 28 | +encodes. The final prediction file (with the majority-vote answer and honest |
| 29 | +K-probe token accounting) is assembled by phase2/build_sc_submission.py. |
| 30 | +""" |
| 31 | + |
| 32 | +import json |
| 33 | +import os |
| 34 | +import re |
| 35 | +import collections |
| 36 | +from typing import Dict, List |
| 37 | + |
| 38 | +from router_inference.router.base_router import BaseRouter |
| 39 | + |
| 40 | +_BOXED = re.compile(r"\\boxed\{+([^{}]*)\}+") |
| 41 | + |
| 42 | + |
| 43 | +def _norm(s: str) -> str: |
| 44 | + return re.sub(r"[^a-z0-9]", "", str(s).lower()) |
| 45 | + |
| 46 | + |
| 47 | +class CruqSCRouter(BaseRouter): |
| 48 | + """Self-consistency cascade: probe the cheap model K times, escalate on disagreement.""" |
| 49 | + |
| 50 | + def __init__(self, router_name: str): |
| 51 | + super().__init__(router_name) |
| 52 | + params = self.config["pipeline_params"] |
| 53 | + self.k = int(params.get("k", 4)) |
| 54 | + self.tau = float(params.get("tau", 0.6)) |
| 55 | + self.probe_model = params.get("probe_model", "qwen/qwen3-235b-a22b-2507") |
| 56 | + self.escalate_model = params.get("escalate_model", "deepseek/deepseek-v4-flash") |
| 57 | + |
| 58 | + here = os.path.dirname(os.path.abspath(__file__)) |
| 59 | + root = os.path.dirname(os.path.dirname(here)) |
| 60 | + # prompt -> global index, so a raw query string can find its cached probes |
| 61 | + self._prompt_to_gi: Dict[str, str] = {} |
| 62 | + for path in ("dataset/router_data.json", "dataset/router_data_10.json"): |
| 63 | + p = os.path.join(root, path) |
| 64 | + if os.path.exists(p): |
| 65 | + for e in json.load(open(p, encoding="utf-8")): |
| 66 | + self._prompt_to_gi[e["prompt_formatted"]] = e["global index"] |
| 67 | + |
| 68 | + # global index -> list of normalized boxed answers from the K probes |
| 69 | + self._samples: Dict[str, List[str]] = collections.defaultdict(list) |
| 70 | + cache = os.path.join(root, "phase2", "data", "qwen_sc_full.jsonl") |
| 71 | + if os.path.exists(cache): |
| 72 | + for line in open(cache, encoding="utf-8"): |
| 73 | + try: |
| 74 | + r = json.loads(line) |
| 75 | + except Exception: |
| 76 | + continue |
| 77 | + if r["s"] < self.k: |
| 78 | + self._samples[r["gi"]].append(_norm(r["boxed"])) |
| 79 | + |
| 80 | + def _get_prediction(self, query: str) -> str: |
| 81 | + gi = self._prompt_to_gi.get(query) |
| 82 | + if gi is None: |
| 83 | + # Unknown query (no cached probes): fall back to the cheap probe model. |
| 84 | + return self.probe_model |
| 85 | + samples = self._samples.get(gi, []) |
| 86 | + boxes = [b for b in samples if b] |
| 87 | + if not boxes: |
| 88 | + # Free-form dataset (no \boxed answer): self-consistency can't apply, so keep |
| 89 | + # the cheap probe model rather than pay to escalate on a signal we don't have. |
| 90 | + return self.probe_model |
| 91 | + top = collections.Counter(boxes).most_common(1)[0][1] |
| 92 | + consistency = top / len(samples) |
| 93 | + return self.probe_model if consistency >= self.tau else self.escalate_model |
0 commit comments