{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":118448,"databundleVersionId":14559231,"isSourceIdPinned":false},{"sourceType":"competition","sourceId":88432,"databundleVersionId":10128463,"isSourceIdPinned":false},{"sourceType":"competition","sourceId":113558,"databundleVersionId":14878066,"isSourceIdPinned":false},{"sourceType":"modelInstanceVersion","sourceId":510361,"databundleVersionId":13306747,"modelInstanceId":404469,"modelId":422375,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nAIMO3 - RANK-1 CHALLENGER  (Final Production Version)\n======================================================\nRoot cause analysis of all errors so far:\n  1. kaggle_evaluation not found    -> sys.path injection            [FIXED]\n  2. SERVER_MISSING_ENDPOINT        -> lazy load inside predict()    [FIXED]\n  3. Mxfp4 vs BitsAndBytesConfig    -> patch config.json before load [FIXED]\n  4. load_in_8bit kwarg not accepted -> removed, use bnb config patch [FIXED]\n  5. CUDA OOM                        -> force 4bit after config patch [FIXED]\n  6. DEVICE=cpu / No GPU             -> GPU must be set in Settings   [INFO]\n\nIMPORTANT - Kaggle Settings required:\n  Settings -> Accelerator -> GPU T4 x2   (NOT None/CPU)\n  Settings -> Internet    -> OFF\n\nThe model gpt-oss-20b uses a custom class GptOssForCausalLM.\nIt stores Mxfp4 quantization in config.json.\nT4 has no Triton -> transformers dequants to bf16 -> OOM.\nFix: patch the AutoConfig object to remove quantization_config\n     BEFORE calling from_pretrained, then load fresh with 4-bit bnb.\n\"\"\"\n\nimport os\nimport re\nimport sys\nimport time\nimport json\nimport threading\nimport traceback\nimport collections\nimport subprocess\n\nos.environ[\"PYTORCH_ALLOC_CONF\"] = \"expandable_segments:True\"\n\nimport torch\nimport polars as pl\n\n# =============================================================================\n# 0.  kaggle_evaluation PATH FIX\n# =============================================================================\ndef _inject_kaggle_eval():\n    for c in [\n        \"/kaggle/input/competitions/ai-mathematical-olympiad-progress-prize-3\",\n        \"/kaggle/input/ai-mathematical-olympiad-progress-prize-3\",\n    ]:\n        if os.path.isdir(os.path.join(c, \"kaggle_evaluation\")):\n            if c not in sys.path:\n                sys.path.insert(0, c)\n            print(f\"[PATH] kaggle_evaluation -> {c}\")\n            return\n    for root, dirs, _ in os.walk(\"/kaggle/input\"):\n        if \"kaggle_evaluation\" in dirs:\n            if root not in sys.path:\n                sys.path.insert(0, root)\n            print(f\"[PATH] kaggle_evaluation -> {root}\")\n            return\n    print(\"[PATH] WARNING: kaggle_evaluation not found!\")\n\n_inject_kaggle_eval()\n\n# =============================================================================\n# 1.  PATH DISCOVERY\n# =============================================================================\ndef _walk_find(name):\n    for root, dirs, files in os.walk(\"/kaggle/input\"):\n        if name in files:\n            return os.path.join(root, name)\n    return None\n\ndef _find_model_dir():\n    known = \"/kaggle/input/models/danielhanchen/gpt-oss-20b/transformers/default/1\"\n    if os.path.isdir(known):\n        return known\n    cfg = _walk_find(\"config.json\")\n    return os.path.dirname(cfg) if cfg else None\n\ndef _find_test_csv():\n    for p in [\n        \"/kaggle/input/competitions/ai-mathematical-olympiad-progress-prize-3/test.csv\",\n        \"/kaggle/input/ai-mathematical-olympiad-progress-prize-3/test.csv\",\n        \"/kaggle/input/test.csv\",\n    ]:\n        if os.path.exists(p):\n            return p\n    found = _walk_find(\"test.csv\")\n    if found:\n        return found\n    raise FileNotFoundError(\"test.csv not found — add competition data\")\n\nMODEL_DIR = _find_model_dir()\nDEVICE    = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"[INIT] MODEL_DIR = {MODEL_DIR}\")\nprint(f\"[INIT] DEVICE    = {DEVICE}\")\n\nif DEVICE == \"cpu\":\n    print(\n        \"\\n[WARNING] No GPU detected!\\n\"\n        \">>> Go to Settings -> Accelerator -> GPU T4 x2\\n\"\n        \"    Then re-run the notebook.\\n\"\n    )\n\n# =============================================================================\n# 2.  PYTHON SANDBOX\n# =============================================================================\n_PREAMBLE = \"\"\"\\\nimport math, re, sys, itertools, functools, collections\nimport numpy as np\ntry:\n    import sympy as sp\n    from sympy import (symbols, solve, simplify, factorint, gcd, lcm,\n                       isprime, nextprime, binomial, factorial,\n                       floor, ceiling, Rational, sqrt, Mod, Integer,\n                       Sum, Product, Matrix)\nexcept ImportError:\n    pass\n\"\"\"\n\ndef _run_code(code, timeout=20):\n    try:\n        r = subprocess.run(\n            [sys.executable, \"-c\", _PREAMBLE + \"\\n\" + code],\n            capture_output=True, text=True, timeout=timeout,\n        )\n        out = r.stdout.strip()\n        return out if out else None\n    except Exception:\n        return None\n\n# =============================================================================\n# 3.  ANSWER EXTRACTION\n# =============================================================================\n_PATS = [\n    r\"\\\\boxed\\{(\\d+)\\}\",\n    r\"\\banswer\\s*=\\s*(\\d+)\",\n    r\"the answer is\\s*:?\\s*(\\d+)\",\n    r\"final answer\\s*:?\\s*(\\d+)\",\n    r\"\\bresult\\s*=\\s*(\\d+)\",\n    r\"=\\s*(\\d+)\\s*$\",\n    r\"(\\d{1,6})\\s*$\",\n]\n\ndef _extract(text):\n    text = str(text)\n    for p in _PATS:\n        m = re.search(p, text, re.IGNORECASE | re.MULTILINE)\n        if m:\n            v = int(m.group(1))\n            if 0 <= v <= 99999:\n                return v\n    for n in reversed(re.findall(r\"\\d+\", text)):\n        v = int(n)\n        if 0 <= v <= 99999:\n            return v\n    return None\n\n# =============================================================================\n# 4.  SYMPY FAST-PATH\n# =============================================================================\ndef _sympy_fast(problem):\n    try:\n        import sympy\n        m = re.search(\n            r\"remainder when\\s+(.+?)\\s+is divided by\\s+(\\d+)\", problem, re.I)\n        if m:\n            return int(sympy.sympify(m.group(1)).evalf()) % int(m.group(2))\n        m2 = re.search(r\"(?:find\\s+)?(.+?)\\s+mod\\s+(\\d+)\", problem, re.I)\n        if m2:\n            return int(sympy.sympify(m2.group(1)).evalf()) % int(m2.group(2))\n    except Exception:\n        pass\n    return None\n\n# =============================================================================\n# 5.  PATCH config.json IN-MEMORY  (core OOM fix)\n#\n#  The model's config.json contains:\n#    \"quantization_config\": {\"quant_type\": \"mxfp4\", ...}\n#\n#  On T4 (no Triton), transformers sees this and:\n#    -> tries mxfp4 -> fails -> falls back to bf16 dequant -> OOM\n#\n#  Fix: load the AutoConfig, delete quantization_config from it,\n#  then pass that patched config + BitsAndBytesConfig 4bit to\n#  from_pretrained. This way bnb 4bit loads cleanly (~10 GB).\n# =============================================================================\ndef _load_patched_model(model_dir):\n    from transformers import AutoTokenizer, AutoModelForCausalLM, AutoConfig\n    from transformers import BitsAndBytesConfig\n\n    print(\"[LOAD] Reading config …\")\n    cfg = AutoConfig.from_pretrained(model_dir, trust_remote_code=True)\n\n    # Remove the Mxfp4 quantization so transformers won't try to use it\n    if hasattr(cfg, \"quantization_config\"):\n        print(f\"[LOAD] Removing stored quantization_config: \"\n              f\"{getattr(cfg, 'quantization_config', None)}\")\n        del cfg.quantization_config\n\n    bnb = BitsAndBytesConfig(\n        load_in_4bit=True,\n        bnb_4bit_quant_type=\"nf4\",\n        bnb_4bit_compute_dtype=torch.float16,\n        bnb_4bit_use_double_quant=True,\n    )\n\n    print(\"[LOAD] Loading tokenizer …\")\n    tok = AutoTokenizer.from_pretrained(model_dir, trust_remote_code=True)\n\n    print(\"[LOAD] Loading model with 4-bit bnb (patched config) …\")\n    mdl = AutoModelForCausalLM.from_pretrained(\n        model_dir,\n        config=cfg,                  # patched — no mxfp4\n        quantization_config=bnb,     # 4-bit bnb instead (~10 GB)\n        device_map=\"auto\",\n        trust_remote_code=True,\n        low_cpu_mem_usage=True,\n    )\n    print(\"[LOAD] Model ready ✓\")\n    return tok, mdl\n\n# =============================================================================\n# 6.  SOLVER\n# =============================================================================\nclass OlympiadSolver:\n    def __init__(self):\n        self._tok  = None\n        self._mdl  = None\n        self._lock = threading.Lock()\n\n    def _lazy_load(self):\n        if self._mdl is not None:\n            return\n        with self._lock:\n            if self._mdl is not None:\n                return\n            if MODEL_DIR is None:\n                print(\"[WARN] No model — dummy mode (answers = 0)\")\n                return\n            if DEVICE == \"cpu\":\n                print(\"[WARN] CPU only — dummy mode (answers = 0)\")\n                return\n            self._tok, self._mdl = _load_patched_model(MODEL_DIR)\n\n    @staticmethod\n    def _tir_prompt(problem):\n        return (\n            \"<|im_start|>system\\n\"\n            \"You are a world-class mathematical olympiad solver. \"\n            \"Use Python/sympy to compute exact integer answers.\\n\"\n            \"<|im_end|>\\n\"\n            \"<|im_start|>user\\n\"\n            f\"{problem}\\n\\n\"\n            \"Write Python code. End with:\\n\"\n            \"  answer = <integer>\\n\"\n            \"  print(answer)\\n\"\n            \"<|im_end|>\\n\"\n            \"<|im_start|>assistant\\n\"\n            \"```python\\n\"\n        )\n\n    @staticmethod\n    def _cot_prompt(problem):\n        return (\n            \"<|im_start|>system\\n\"\n            \"You are a brilliant mathematician at IMO level.\\n\"\n            \"<|im_end|>\\n\"\n            \"<|im_start|>user\\n\"\n            f\"{problem}\\n\\n\"\n            \"Think step by step. Answer is a non-negative integer <= 99999.\\n\"\n            \"Last line: The answer is: <integer>\\n\"\n            \"<|im_end|>\\n\"\n            \"<|im_start|>assistant\\n\"\n        )\n\n    def _gen(self, prompt, temperature):\n        enc = self._tok(prompt, return_tensors=\"pt\").to(DEVICE)\n        with torch.no_grad():\n            out = self._mdl.generate(\n                **enc,\n                max_new_tokens=1536,\n                temperature=max(float(temperature), 0.01),\n                do_sample=(temperature > 0.02),\n                pad_token_id=self._tok.eos_token_id,\n                eos_token_id=self._tok.eos_token_id,\n            )\n        new = out[0][enc[\"input_ids\"].shape[1]:]\n        return self._tok.decode(new, skip_special_tokens=True)\n\n    def solve(self, problem, n_samples=8, time_limit=480.0):\n        self._lazy_load()\n\n        fast = _sympy_fast(problem)\n        if fast is not None:\n            print(f\"[SYMPY] {fast}\")\n            return fast\n\n        if self._mdl is None:\n            return 0\n\n        t0 = time.time()\n        all_votes  = []\n        code_votes = []\n\n        for i in range(n_samples):\n            if time.time() - t0 > time_limit:\n                print(f\"[TIME] stopped at sample {i}\")\n                break\n\n            if i % 2 == 0:\n                prompt = self._tir_prompt(problem)\n                temp   = 0.05 if i == 0 else round(0.3 + 0.1 * (i // 2), 2)\n            else:\n                prompt = self._cot_prompt(problem)\n                temp   = round(0.5 + 0.1 * (i // 2), 2)\n\n            try:\n                resp = self._gen(prompt, temp)\n            except Exception as e:\n                print(f\"[GEN-{i}] {e}\")\n                continue\n\n            cb = re.search(r\"```python\\n?(.*?)(?:```|$)\", resp, re.DOTALL)\n            if cb is None:\n                cb = re.search(r\"((?:.*\\n)*?answer\\s*=\\s*\\d+.*)\", resp)\n            if cb:\n                out = _run_code(cb.group(1))\n                v   = _extract(out) if out else None\n                if v is not None:\n                    print(f\"[TIR-{i}] temp={temp} -> {v}\")\n                    code_votes.append(v)\n                    all_votes.append(v)\n                    continue\n\n            v = _extract(resp)\n            if v is not None:\n                print(f\"[TXT-{i}] temp={temp} -> {v}\")\n                all_votes.append(v)\n\n        if not all_votes:\n            return 0\n\n        W = collections.defaultdict(float)\n        code_set = set(code_votes)\n        for v in all_votes:\n            W[v] += 2.0 if v in code_set else 1.0\n        keys = list(W.keys())\n        for a in keys:\n            for b in keys:\n                if a != b and abs(a - b) <= 1:\n                    W[a] += 0.25 * W[b]\n        best = max(W, key=lambda k: W[k])\n        print(f\"[VOTE] {dict(W)} -> {best}\")\n        return best\n\n# =============================================================================\n# 7.  GLOBAL SOLVER\n# =============================================================================\n_solver = OlympiadSolver()\n\n# =============================================================================\n# 8.  KAGGLE IMPORTS\n# =============================================================================\nimport kaggle_evaluation.aimo_3_inference_server  # noqa: E402\n\n# =============================================================================\n# 9.  PREDICT\n# =============================================================================\ndef predict(id_: pl.Series, problem: pl.Series) -> pl.DataFrame:\n    pid   = id_.item(0)\n    ptext = problem.item(0)\n    print(f\"\\n{'='*60}\\n[Q] id={pid}\\n{ptext[:220]}\\n{'='*60}\")\n    try:\n        ans = _solver.solve(ptext, n_samples=8, time_limit=480.0)\n    except Exception:\n        traceback.print_exc()\n        ans = 0\n    print(f\"[A] id={pid} -> {ans}\")\n    return pl.DataFrame({\"id\": pid, \"answer\": int(ans)})\n\n# =============================================================================\n# 10. ENTRY POINT\n# =============================================================================\ninference_server = kaggle_evaluation.aimo_3_inference_server.AIMO3InferenceServer(\n    predict\n)\n\nif os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"):\n    inference_server.serve()\nelse:\n    try:\n        test_path = _find_test_csv()\n        print(f\"[RUN] test: {test_path}\")\n        inference_server.run_local_gateway((test_path,))\n    except FileNotFoundError as e:\n        print(f\"[ERROR] {e}\")","metadata":{"_uuid":"0e6bdc2c-2266-4ac4-8bd1-ae15aae9414a","_cell_guid":"21ba8757-674c-4522-87a7-c89bd3dc6e65","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}