{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"nvidiaL4","dataSources":[{"sourceId":84795,"databundleVersionId":10934030,"sourceType":"competition"},{"sourceId":221096520,"sourceType":"kernelVersion"},{"sourceId":236932,"sourceType":"modelInstanceVersion","modelInstanceId":202348,"modelId":224071}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\n# https://www.kaggle.com/competitions/ai-mathematical-olympiad-progress-prize-2/discussion/560682#3113134\nos.environ[\"TRITON_PTXAS_PATH\"] = \"/usr/local/cuda/bin/ptxas\"","metadata":{"_cell_guid":"eec2b282-90f3-4965-a943-54cffbeb7a5b","_uuid":"38b86f00-408f-4577-b91e-890da62fd434","collapsed":false,"execution":{"iopub.execute_input":"2025-02-18T05:45:10.317457Z","iopub.status.busy":"2025-02-18T05:45:10.317295Z","iopub.status.idle":"2025-02-18T05:45:10.321715Z","shell.execute_reply":"2025-02-18T05:45:10.320472Z","shell.execute_reply.started":"2025-02-18T05:45:10.317439Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import io\nimport time\nimport shutil\nimport subprocess \nimport pandas as pd\nimport polars as pl\nfrom concurrent.futures import ThreadPoolExecutor\nfrom multiprocessing.pool import ThreadPool\n\nimport kaggle_evaluation.konwinski_prize_inference_server\nfrom typing import List, Tuple, Dict, Optional\n\nstart_time = time.time()","metadata":{"_cell_guid":"2512bd4d-9caf-4c50-9706-0377f47811ae","_kg_hide-output":true,"_uuid":"e9ab5832-846d-4fb5-b21b-532f5962301d","collapsed":false,"execution":{"iopub.execute_input":"2025-02-18T05:45:10.322694Z","iopub.status.busy":"2025-02-18T05:45:10.322482Z","iopub.status.idle":"2025-02-18T05:45:29.670731Z","shell.execute_reply":"2025-02-18T05:45:29.66948Z","shell.execute_reply.started":"2025-02-18T05:45:10.322673Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"instance_count: Optional[int] = None\n\n\ndef get_number_of_instances(num_instances: int) -> None:\n    \"\"\"The very first message from the gateway will be the total number of instances to be served.\n    You don't need to edit this function.\n    \"\"\"\n    global instance_count\n    instance_count = num_instances","metadata":{"_cell_guid":"7fce3650-d66d-4189-ac61-51f4e07ac289","_uuid":"f251dce7-62ae-4856-9c05-f5ffe2938eaf","collapsed":false,"execution":{"iopub.execute_input":"2025-02-18T05:45:29.672776Z","iopub.status.busy":"2025-02-18T05:45:29.671707Z","iopub.status.idle":"2025-02-18T05:45:29.677844Z","shell.execute_reply":"2025-02-18T05:45:29.676609Z","shell.execute_reply.started":"2025-02-18T05:45:29.672748Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.011949,"end_time":"2024-12-11T03:22:08.838279","exception":false,"start_time":"2024-12-11T03:22:08.82633","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from vllm import LLM, SamplingParams, RequestOutput\nimport warnings\n\nwarnings.simplefilter(\"ignore\")\n\nos.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0,1,2,3\"\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\nif os.getenv(\"KAGGLE_KERNEL_RUN_TYPE\") or os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"):\n    llm_model_pth: str = (\n        \"/kaggle/input/deepseek-r1/transformers/deepseek-r1-distill-qwen-32b-awq/1\"\n    )\nelse:\n    llm_model_pth: str = \"/root/volume/KirillR/QwQ-32B-Preview-AWQ\"\n\nBATCH_SIZE: int = 6\nVALIDATION_COPY_COUNT: int = 1\nMAX_TOKENS: int = 4096\n\nMAX_NUM_SEQS: int = 6\nMAX_MODEL_LEN: int = 32_768\n\nllm: LLM = LLM(\n    llm_model_pth,\n    max_num_seqs=MAX_NUM_SEQS,  # Maximum number of sequences per iteration. Default is 256\n    max_model_len=MAX_MODEL_LEN,  # Model context length\n    trust_remote_code=True,  # Trust remote code (e.g., from HuggingFace) when downloading the model and tokenizer\n    tensor_parallel_size=4,  # The number of GPUs to use for distributed execution with tensor parallelism\n    gpu_memory_utilization=0.95,  # The ratio (between 0 and 1) of GPU memory to reserve for the model\n    seed=2024,\n)","metadata":{"_cell_guid":"c863b090-e04a-441e-9add-bddae0863ef9","_kg_hide-output":true,"_uuid":"eca28397-2a40-4265-b36e-dfa9d80eb2b1","collapsed":false,"execution":{"iopub.execute_input":"2025-02-18T05:45:29.680772Z","iopub.status.busy":"2025-02-18T05:45:29.680387Z","iopub.status.idle":"2025-02-18T05:46:58.41742Z","shell.execute_reply":"2025-02-18T05:46:58.41444Z","shell.execute_reply.started":"2025-02-18T05:45:29.680737Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tokenizer = llm.get_tokenizer()\n\n\ndef count_tokens(text: str) -> int:\n    return len(tokenizer.encode(text))","metadata":{"_cell_guid":"99e40ae3-6474-4b05-a969-0b48de303ba6","_uuid":"5a982fe4-e4e2-466d-b5bf-4839ec8804c8","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.418472Z","iopub.status.idle":"2025-02-18T05:46:58.418853Z","shell.execute_reply":"2025-02-18T05:46:58.418713Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\n\ndef stringify_directory(directory: str) -> str:\n    full_paths: List[str] = []\n\n    for root, dirs, files in os.walk(directory):\n        for file in files:\n            full_path: str = os.path.join(root, file)\n            full_paths.append(full_path)\n    return \"\\n\".join(full_paths)","metadata":{"_cell_guid":"44d40acd-8428-4813-ba5d-62f625309dd4","_uuid":"43546fee-8473-4f0d-80ec-46a86369426b","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.422805Z","iopub.status.idle":"2025-02-18T05:46:58.423335Z","shell.execute_reply":"2025-02-18T05:46:58.423121Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import re\n\n\ndef extract_file_query(xml_content: str) -> Dict[str, List[str]]:\n    import xml.etree.ElementTree as ET\n\n    # Prepare a data structure to collect results\n    parsed_data: Dict[str, List[str]] = {}\n    pattern: str = r\"<root>(.*?)</root>\"\n    matches: List[str] = re.findall(pattern, xml_content, re.DOTALL)\n\n    for match in matches:\n        try:\n            # Parse the XML\n            root = ET.fromstring(\"<root>\" + match + \"</root>\")\n\n            # Find all <entry> elements\n            for entry in root.findall(\"entry\"):\n                # Extract the <filepath> text\n                filepath = entry.find(\"filepath\")\n                filepath_text: Optional[str] = (\n                    filepath.text.strip()\n                    if filepath is not None and filepath.text is not None\n                    else None\n                )\n\n                # Locate <strings_to_search> container\n                strings_container = entry.find(\"strings_to_search\")\n\n                # Gather each <string_to_search> text\n                search_strings: List[str] = []\n                if strings_container is not None:\n                    for s in strings_container.findall(\"string_to_search\"):\n                        if s.text is not None:\n                            search_strings.append(s.text.strip())\n\n                # Store in a dictionary: { filepath: [search_strings...] }\n                parsed_data[filepath_text] = search_strings  # type: ignore\n        except:\n            print(\"Error parsing output\")\n            print(xml_content)\n            return {}\n\n    return parsed_data","metadata":{"_cell_guid":"50b46124-4395-4783-b523-f11f4b6228bf","_uuid":"47b13790-f9e8-47cd-8b78-c93567e04b56","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.424233Z","iopub.status.idle":"2025-02-18T05:46:58.42457Z","shell.execute_reply":"2025-02-18T05:46:58.424438Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import re\n\nreading_prompt: str = (\n    \"\"\"\nYou will be implementing a git diff patch to solve an issue with the code repository.\nYou will first need to select files in the file directory.\n\nThis is the problem statement.\n\n{problem_statement}\n\nThis is the file directory\n\n<directory>\n{directory_string}\n</directory>\n\nWhich files should be inspected so that we can solve the problem?\nWhen we inspect each file, what strings should be searched?\n\nReturn the strings to search in this format\n\n(explanation)\n\n<root>\n    <entry>\n        <filepath>filepath</filepath>\n        <strings_to_search>\n            <string_to_search>string_to_search</string_to_search>\n            ...\n            <string_to_search>string_to_search</string_to_search>\n        </strings_to_search>\n    </entry>\n    <entry>\n        <filepath>filepath</filepath>\n        <strings_to_search>\n            <string_to_search>string_to_search</string_to_search>\n            ...\n            <string_to_search>string_to_search</string_to_search>\n        </strings_to_search>\n    </entry>\n    ...\n</root>\n...\n\nNotes:\n- Make sure to encode each entry between <root> and </root>\n- Return the FULL filepath - exactly as specified in <directory> and </directory>\n    - Example: <filepath>repo/path/to/directory/file.py</filepath>\n- If you are searching for a word instead of a substring, maybe add spaces or brackets before and after the string\n    - For example, if you are searching for uses of the function `calculate`, use ` calculate(` as the search string instead of `calculate`\n- Prefer searching longer strings\n    - Avoid searching for strings that might appear in many parts of the codebase\n- Search the test files as well to understand the feature behavior\n    - Also search for the relevant function calls in the test files\n\"\"\".strip()\n)\n\n\ndef get_selection_query(\n    directory_string: str, problem_statement: str\n) -> Tuple[List[str], List[Dict[str, List[str]]]]:\n    sampling_params: SamplingParams = SamplingParams(\n        temperature=0.6,  # randomness of the sampling\n        min_p=0.01,\n        skip_special_tokens=True,  # Whether to skip special tokens in the output\n        max_tokens=MAX_TOKENS,\n    )\n\n    list_of_messages: List[List[Dict[str, str]]] = [\n        [\n            {\n                \"role\": \"user\",\n                \"content\": reading_prompt.format(\n                    problem_statement=problem_statement[:20_000],\n                    directory_string=directory_string[:30_000],\n                ),\n            },\n        ]\n        for _ in range(BATCH_SIZE)\n    ]\n\n    prompt_texts: List[str] = [\n        (\n            tokenizer.apply_chat_template(\n                conversation=messages, tokenize=False, add_generation_prompt=True\n            )  # type: ignore\n        )\n        + \"<think>\\n\"\n        for messages in list_of_messages\n    ]\n    # print(prompt_texts)\n\n    print(\"get_selection_query\", [count_tokens(text) for text in prompt_texts])\n    request_outputs: list[RequestOutput] = llm.generate(\n        prompt_texts, sampling_params=sampling_params\n    )\n    if not request_outputs:\n        return [], []\n    response_texts: List[str] = [\n        request_output.outputs[0].text for request_output in request_outputs\n    ]\n    print(\"get_selection_query\", [count_tokens(text) for text in response_texts])\n\n    completion_texts = [\n        prompt_text + response_text\n        for prompt_text, response_text in zip(prompt_texts, response_texts)\n    ]\n    file_queries: List[Dict[str, List[str]]] = [\n        extract_file_query(response_text) for response_text in response_texts\n    ]\n    return completion_texts, file_queries","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"REPO_PATH: str = \"repo\"\n\n\ndef fetch_file_contents(\n    files_to_search: Dict[str, List[str]], context_lines: int = 12, max_gap: int = 0\n) -> str:\n    from io import StringIO\n    from typing import Tuple\n\n    def find_lines_in_files_with_context(\n        search_map: Dict[str, List[str]], context_lines: int = context_lines\n    ) -> List[List[List[Tuple[int, str]]]]:\n        \"\"\"\n        Given a dictionary mapping file paths to a list of search terms,\n        open each file and gather *snippets* of lines that contain any\n        of those search terms, including 'context_lines' before and after.\n\n        Returns a list of lists:\n        [\n          [  # For file1\n             [ (line_number, text), (line_number, text), ... ],\n             [ ... ],\n          ],\n          [  # For file2\n             ...\n          ],\n          ...\n        ]\n        \"\"\"\n        all_matches_per_file: List[List[List[Tuple[int, str]]]] = []\n\n        for path, terms in search_map.items():\n            if not os.path.isfile(path):\n                # If the file is not found, record an empty list\n                all_matches_per_file.append([])\n                continue\n\n            with open(path, \"r\", encoding=\"utf-8\", errors=\"replace\") as f:\n                lines = f.readlines()\n\n            file_snippets: List[List[Tuple[int, str]]] = []\n            num_lines: int = len(lines)\n\n            for i, line in enumerate(lines, start=1):\n                if any(t in line for t in terms):\n                    start_idx: int = max(1, i - context_lines)\n                    end_idx: int = min(num_lines, i + context_lines)\n                    snippet: List[Tuple[int, str]] = []\n                    for snippet_no in range(start_idx, end_idx + 1):\n                        text_content: str = lines[snippet_no - 1].rstrip(\"\\n\")\n                        snippet.append((snippet_no, text_content))\n                    file_snippets.append(snippet)\n\n            all_matches_per_file.append(file_snippets)\n\n        return all_matches_per_file\n\n    # ---------------------------------------------------------\n    # 3. MERGE OVERLAPPING/ADJACENT SNIPPETS\n    # ---------------------------------------------------------\n\n    def merge_file_snippets(\n        file_snippets: List[List[Tuple[int, str]]], gap: int = 0\n    ) -> List[List[Tuple[int, str]]]:\n        \"\"\"\n        Merge overlapping or nearly adjacent snippets in a single file’s snippet list.\n        \"\"\"\n        intervals: List[Tuple[int, int, List[Tuple[int, str]]]] = []\n        for snippet in file_snippets:\n            if snippet:\n                start_line: int = snippet[0][0]\n                end_line: int = snippet[-1][0]\n                intervals.append((start_line, end_line, snippet))\n\n        intervals.sort(key=lambda x: x[0])  # sort by start line\n\n        merged: List[Tuple[int, int, List[Tuple[int, str]]]] = []\n        for start, end, snippet in intervals:\n            if not merged:\n                merged.append((start, end, snippet))\n                continue\n\n            prev_start, prev_end, prev_snippet = merged[-1]\n            if start <= prev_end + gap:\n                new_end: int = max(end, prev_end)\n                combined_dict: Dict[int, str] = {}\n                for ln, txt in prev_snippet:\n                    combined_dict[ln] = txt\n                for ln, txt in snippet:\n                    combined_dict[ln] = txt\n                merged_snippet: List[Tuple[int, str]] = [\n                    (ln, combined_dict[ln]) for ln in sorted(combined_dict)\n                ]\n                merged[-1] = (prev_start, new_end, merged_snippet)\n            else:\n                merged.append((start, end, snippet))\n\n        # Extract just the merged snippet portion\n        return [x[2] for x in merged]\n\n    def merge_all_snippets(\n        all_files_snips: List[List[List[Tuple[int, str]]]], gap: int = 0\n    ) -> List[List[List[Tuple[int, str]]]]:\n        \"\"\"\n        Merge snippet blocks within each file.\n        all_files_snips is a list-of-lists:\n          [\n            [ snippetA, snippetB, ... ],  # file 1\n            [ snippetC, snippetD, ... ],  # file 2\n          ]\n        \"\"\"\n        merged: List[List[List[Tuple[int, str]]]] = []\n        for snips in all_files_snips:\n            merged.append(merge_file_snippets(snips, gap=gap))\n        return merged\n\n    # ---------------------------------------------------------\n    # 4. RUN LOGIC: generate files, search, merge, and BUILD A STRING\n    # ---------------------------------------------------------\n\n    has_any_matches: bool = False\n\n    # 1) Gather snippets around each match\n    context_snippets: List[List[List[Tuple[int, str]]]] = (\n        find_lines_in_files_with_context(files_to_search, context_lines=context_lines)\n    )\n\n    # 2) Merge overlapping snippets\n    merged_snips: List[List[List[Tuple[int, str]]]] = merge_all_snippets(\n        context_snippets, gap=max_gap\n    )\n\n    # 3) Build a string (instead of printing)\n    output = StringIO()\n\n    # Header\n    output.write(\"Sample files created successfully.\\n\\n\")\n    output.write(\"Search Results (by file, merging any overlapping context):\\n\\n\")\n\n    # For each file\n    for (filepath, terms), snippet_list in zip(files_to_search.items(), merged_snips):\n        output.write(f\"[file name]: {filepath[len(REPO_PATH) + 1:]}\\n\")\n        terms_searched_as_str = \"\\n\".join(terms)\n        output.write(f\"[terms searched]:\\n{terms_searched_as_str}\\n\")\n        output.write(\"[file content begin]\\n\")\n        if not snippet_list:\n            output.write(\"  No matches found.\\n\")\n        else:\n            has_any_matches = True\n            for snippet_idx, snippet in enumerate(snippet_list, start=1):\n                snippet_start: int = snippet[0][0]\n                snippet_end: int = snippet[-1][0]\n                output.write(\n                    f\"\\nMatch #{snippet_idx}, lines {snippet_start} to {snippet_end}:\\n\"\n                )\n                for line_no, text in snippet:\n                    output.write(f\"  {line_no:3d} | {text}\\n\")\n                output.write(\"\\n\")\n        output.write(\"[file content end]\\n\\n\")\n\n    file_content_string: str = output.getvalue()\n\n    if has_any_matches:\n        return file_content_string\n    return \"\"","metadata":{"_cell_guid":"2e6aebd6-10e8-4769-98dc-4e89aaf35b41","_uuid":"f8ae44f3-2595-4c78-b58a-6551c0fec3a8","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.426342Z","iopub.status.idle":"2025-02-18T05:46:58.426591Z","shell.execute_reply":"2025-02-18T05:46:58.426491Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import re\n\n# --- Helper Functions - Patch Extraction and Test Outcome ---\n\ndef extract_patch_string(text: str) -> Optional[str]:\n    pattern: str = r\"\\n```diff\\n(.*?)\\n```\"\n    matches: List[str] = re.findall(pattern, text, re.DOTALL)\n    if not matches:\n        return None\n    return matches[-1] + \"\\n\"\n\ndef get_unit_test_outcome(repo_path: str) -> List[str]: # Changed return type to List[str]\n    \"\"\"\n    Simplified test running logic.  Replace with the *exact* logic\n    used in the evaluation (if available).\n    \"\"\"\n    # Simulate running unit tests (replace with actual test execution)\n    # This is crucial for MCTS to work effectively\n    cmd = f\"python -m unittest discover -s {repo_path}\"\n    try:\n        result = subprocess.run(cmd, shell=True, check=True, capture_output=True, text=True, timeout=60)\n        # If successful, all tests passed (crude approximation)\n        return [\"PASSED\"]  # Return a list to mimic the original logic\n    except subprocess.CalledProcessError as e:\n        # If there was an error, some tests failed\n        # (Parse the output to identify individual failed tests, if possible)\n        print(f\"Tests failed: {e.stderr}\")\n        return [\"FAILED\"]","metadata":{"_cell_guid":"a2564919-7162-4711-a5c6-92077938254f","_uuid":"5dcae870-2699-4110-956b-121c8687bfa9","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.427173Z","iopub.status.idle":"2025-02-18T05:46:58.427452Z","shell.execute_reply":"2025-02-18T05:46:58.427352Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\n\nverifying_prompt: str = (\n    \"\"\"\nThis is the problem statement.\n\n{problem_statement}\n\nThese are the files that is thought to be relevant, which may not be complete.\n\n{file_content_string}\n\nThis is the proposed patch to fix the problem.\n\n{patch_string}\n\nEvaluate whether the patch works\n- The patch fully fixes the problem described in the problem statement.\n- The patch does not cause side effects and make any other tests fail.\n\nEnd your response with exactly either of\n- <label>Yes</label>, this fixes the problem.\n- <label>No</label>, this does not fix the problem.\n\nReminder\n- Only evaluate, do not provide suggestion on how to fix.\n- Remember to write exactly either of <label>Yes</label> or <label>No</label> in the last line\n\"\"\".strip()\n)\n\n\ndef is_valid_patch_format(patch_string: str) -> bool:\n    \"\"\"\n    A quick check to confirm if a patch could be valid.\n    \"\"\"\n    if not(isinstance(patch_string, str)):\n        return False\n    try:\n        patch_set = unidiff.PatchSet(patch_string)\n        if len(patch_set) == 0:\n            return False\n    except Exception:\n        return False\n    return True\n\n\ndef patch_dry_run_succeeds(patch_string: str, repo_path: str = REPO_PATH, timeout: int = 60) -> bool:\n    \"\"\"\n    A robust check if the patch will proceed without any errors.\n    Should be run after `is_valid_patch_format()`: the patch\n    command can hang if the inputs are sufficiently invalid.\n\n    Args:\n        patch_path: Path to a file containing the patch.\n        repo_path: Path to the directory to be patched.\n        timeout: Number of seconds before the dry run will be cancelled.\n    \"\"\"\n    with open(\"patch.txt\", \"w\") as f:\n        f.write(patch_string)\n    patch_path = \"/kaggle/working/patch.txt\"\n\n    cmd = f\"patch --quiet --dry-run -p1 -i {patch_path} -d {repo_path}\"\n    try:\n        subprocess.run(cmd, shell=True, check=True, timeout=timeout)\n        return True\n    except subprocess.CalledProcessError:\n        return False\n\n\ndef get_verification(\n    problem_statement: str,\n    file_content_strings: List[str],\n    patch_strings: List[Optional[str]],\n    repo_path: str,\n) -> Tuple[List[List[str]], List[List[bool]]]:\n    assert len(file_content_strings) == len(patch_strings)\n    sampling_params: SamplingParams = SamplingParams(\n        temperature=0.3,  # randomness of the sampling\n        min_p=0.01,\n        skip_special_tokens=True,  # Whether to skip special tokens in the output\n        max_tokens=MAX_TOKENS,\n    )\n\n    inference_idx_to_input_idx: list[int] = [\n        input_idx\n        for _ in range(VALIDATION_COPY_COUNT)\n        for input_idx, patch_string in enumerate(patch_strings)\n        if patch_string is not None and is_valid_patch_format(patch_string) # and patch_dry_run_succeeds(patch_string, repo_path)\n    ]\n    print(inference_idx_to_input_idx)\n\n    list_of_messages: List[List[Dict[str, str]]] = [\n        [\n            {\n                \"role\": \"user\",\n                \"content\": verifying_prompt.format(\n                    problem_statement=problem_statement[:20_000],\n                    file_content_string=file_content_strings[input_idx][:30_000],\n                    patch_string=patch_strings[input_idx],\n                ),\n            },\n        ]\n        for input_idx in inference_idx_to_input_idx\n    ]\n\n    prompt_texts: List[str] = [\n        (\n            tokenizer.apply_chat_template(\n                conversation=messages, tokenize=False, add_generation_prompt=True\n            )  # type: ignore\n        )\n        + \"<think>\\n\"\n        for messages in list_of_messages\n    ]\n    # print(prompt_texts)\n\n    print(\"get_verification\", [count_tokens(text) for text in prompt_texts])\n    request_outputs: list[RequestOutput] = llm.generate(\n        prompt_texts, sampling_params=sampling_params\n    )\n    response_texts: List[str] = [\n        request_output.outputs[0].text for request_output in request_outputs\n    ]\n    print(\"get_verification\", [count_tokens(text) for text in response_texts])\n\n    completion_texts = [\n        prompt_text + response_text\n        for prompt_text, response_text in zip(prompt_texts, response_texts)\n    ]\n    judgments_flattened: List[bool] = [\n        \"<label>Yes</label>\" in response_text for response_text in response_texts\n    ]\n    print(judgments_flattened)\n\n    judgments_aggregated: List[List[bool]] = [[] for _ in file_content_strings]\n    completion_text_aggregated: List[List[str]] = [[] for _ in patch_strings]\n    for inference_idx, (completion_text, judgement) in enumerate(\n        zip(completion_texts, judgments_flattened)\n    ):\n        input_idx = inference_idx_to_input_idx[inference_idx]\n        completion_text_aggregated[input_idx].append(completion_text)\n        judgments_aggregated[input_idx].append(judgement)\n    print(judgments_aggregated)\n\n    return completion_text_aggregated, judgments_aggregated","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- MCTS Node Class ---\nimport math # Import math for sqrt and log\n\ndef get_reward(self, repo_path: str) -> float:\n    \"\"\"More detailed reward function.\"\"\"\n    unit_test_outcomes = get_unit_test_outcome(repo_path) # Assuming this returns a list of test outcomes now\n    reward = 0.0\n    num_passed_tests = 0\n    num_failed_tests = 0\n    num_errors = 0\n\n    for outcome in unit_test_outcomes:\n        if outcome == \"PASSED\":\n            reward += 2.0\n            num_passed_tests += 1\n        elif outcome == \"FAILED\":\n            reward -= 3.0\n            num_failed_tests += 1\n        elif outcome == \"ERROR\": # Handle ERROR outcomes\n            reward -= 1.5\n            num_errors += 1\n        elif outcome == \"SKIPPED\": # Handle SKIPPED outcomes\n            pass\n\n    if num_passed_tests == 0 and num_failed_tests > 0:\n        reward -= 1.0\n\n    if self.patch_history and not patch_dry_run_succeeds(self.patch_history[-1], repo_path):\n        reward -= 10.0 # Increased penalty for invalid patch\n\n    return reward\n\n\nclass MCTSNode: # Class definition remains the same as previously optimized\n    def __init__(self, repo_state: str, patch_history: List[str], problem_statement: str, parent=None, selected_files_content: str = \"\", visit_count=0, total_value=0): # Added visit_count and total_value for progressive widening and node reuse\n        self.repo_state = repo_state\n        self.patch_history = patch_history\n        self.problem_statement = problem_statement\n        self.parent = parent\n        self.children: List[MCTSNode] = []\n        self.visits = visit_count # Renamed from visits to visit_count\n        self.score = total_value # Renamed from score to total_value\n        self.untried_actions: List[str] = []\n        self.selected_files_content = selected_files_content\n\n    def generate_actions(self, num_actions=20): # num_actions increased to 20\n        patching_prompt = \"\"\"You are a world-class code patching expert. Your goal is to create the best git diff patch for a given problem... (rest of prompt remains same for CoT)\"\"\".strip() # Chain-of-Thought Prompt\n\n        prompt_texts = [patching_prompt.format(problem_statement=self.problem_statement, code=self.selected_files_content) for _ in range(num_actions)]\n        sampling_params = SamplingParams(\n            temperature=1.0, # Increased temperature for more creative patches\n            top_p=0.95,      # High top_p for broader sampling\n            min_p=0.005,     # Reduced min_p to allow less common tokens\n            skip_special_tokens=True,\n            max_tokens=1024, # Reduced max_tokens for faster generation\n        ) # Optimized sampling params\n        request_outputs: list[RequestOutput] = llm.generate(prompt_texts, sampling_params=sampling_params)\n        response_texts: List[str] = [request_output.outputs[0].text for request_output in request_outputs]\n        patch_strings: List[str] = [extract_patch_string(text) for text in response_texts if extract_patch_string(text) is not None]\n\n        filtered_patches = []\n        for patch_string in patch_strings:\n            if patch_string is None:\n                continue\n            if not is_valid_patch_format(patch_string):\n                continue\n            if patch_string.count('\\n') < 3: # Filter out short patches\n                continue\n            filtered_patches.append(patch_string)\n\n        import random # Import random inside function for clarity\n        if len(filtered_patches) > num_actions:\n            self.untried_actions = random.sample(filtered_patches, num_actions)\n        else:\n            self.untried_actions = filtered_patches\n\n\n    def apply_patch(self, patch_string: str) -> str: # Apply patch logic remains same\n        with open(\"patch.txt\", \"w\") as f:\n            f.write(patch_string)\n        patch_path = \"/kaggle/working/patch.txt\"\n\n        with open(\"repo_content.txt\", \"w\") as f:\n            f.write(self.repo_state)\n        content_file_path = \"/kaggle/working/repo_content.txt\"\n\n        cmd = f\"patch < {patch_path} < {content_file_path}\"\n        try:\n            result = subprocess.run(cmd, shell=True, check=True, capture_output=True, text=True, timeout=60)\n            return result.stdout\n        except subprocess.CalledProcessError as e:\n            print(f\"Patch application failed: {e.stderr}\")\n            return self.repo_state\n\n\n    def get_reward(self, repo_path: str) -> float:\n        \"\"\"More detailed reward function.\"\"\"\n        unit_test_outcomes = get_unit_test_outcome(repo_path) # Assuming this returns a list of test outcomes now\n        reward = 0.0\n        num_passed_tests = 0\n        num_failed_tests = 0\n        num_errors = 0\n\n        for outcome in unit_test_outcomes:\n            if outcome == \"PASSED\":\n                reward += 2.0\n                num_passed_tests += 1\n            elif outcome == \"FAILED\":\n                reward -= 3.0\n                num_failed_tests += 1\n            elif outcome == \"ERROR\":\n                reward -= 1.5\n                num_errors += 1\n            elif outcome == \"SKIPPED\":\n                pass\n\n        if num_passed_tests == 0 and num_failed_tests > 0:\n            reward -= 1.0\n\n        if self.patch_history and not patch_dry_run_succeeds(self.patch_history[-1], repo_path):\n            reward -= 10.0 # Increased penalty for invalid patch\n\n        return reward\n\n    def is_terminal(self) -> bool: # Terminal condition remains same\n        if len(self.patch_history) >= 5:\n            return True\n        return False","metadata":{"_cell_guid":"2000b5e2-13f3-4b86-8743-4aff8a950f2c","_uuid":"3229d17a-0142-4003-8b3c-ba42ae72d4d9","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.429003Z","iopub.status.idle":"2025-02-18T05:46:58.4293Z","shell.execute_reply":"2025-02-18T05:46:58.429168Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- MCTS Algorithm (mcts function) ---\n\ndef mcts(root: MCTSNode, repo_path: str, iterations: int = 100) -> str: # Increased iterations to 100\n    C = 2.0 # Exploration constant\n\n    executor = ThreadPoolExecutor(max_workers=8) # Parallel test execution\n    \n    for _ in range(iterations):\n        node = root\n        # --- Selection using UCT ---\n        while node.untried_actions == [] and node.children != []:\n            best_child = None\n            best_uct = -float(\"inf\")\n            for child in node.children:\n                uct = child.score / child.visits + C * (math.log(node.visits) / child.visits) ** 0.5\n                if uct > best_uct:\n                    best_child = child\n                    best_uct = uct\n            node = best_child  # type: ignore\n\n        # --- Expansion ---\n        if node.untried_actions == [] and not node.is_terminal():\n            if node.visits < 10: # Progressive widening: Explore more if node is not visited much\n                 node.generate_actions(num_actions=5) # Generate fewer actions initially\n            else:\n                node.generate_actions(num_actions=20) # Generate more actions for well-visited nodes\n\n        # --- Simulation ---\n        if node.untried_actions:\n            action = node.untried_actions.pop()\n\n            new_repo_state = node.apply_patch(action)\n            new_patch_history = node.patch_history + [action]\n\n            child = MCTSNode(\n                repo_state=new_repo_state,\n                patch_history=new_patch_history,\n                problem_statement=node.problem_statement,\n                parent=node,\n            )\n            node.children.append(child)\n            node = child\n\n            # Use parallel execution for get_reward (test running)\n            future_reward = executor.submit(child.get_reward, repo_path) # Submit test execution to thread pool\n            reward = future_reward.result() # Wait for test execution to complete and get reward\n        else:\n            #If no untried actions, but not terminal, still do a rollout (exploration)\n            if not node.is_terminal() and node.children: # Check if children exist before attempting rollout\n              node = random.choice(node.children) # Basic rollout: explore existing children\n              reward = node.get_reward(repo_path) # Get reward from rolled-out node\n            else:\n              reward = node.get_reward(repo_path) # Or get reward directly if terminal or no children\n\n\n        # --- Backpropagation ---\n        while node is not None:\n            node.visits += 1\n            node.score += reward\n            node = node.parent\n\n    executor.shutdown(wait=False) # Shutdown thread pool\n\n    # --- Action Choice ---\n    best_child = None\n    best_score = -float('inf')\n    for child in root.children:\n        if child.score > best_score:\n            best_child = child\n            best_score = child.score\n\n    if best_child:\n        return best_child.patch_history[-1]\n    return None # Return None if no patch found","metadata":{"_cell_guid":"b7e7aec9-62b4-4dc2-af35-b8f01bdcff73","_uuid":"5f229582-95fc-48d2-ac8e-27e71fd5e4c9","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.43081Z","iopub.status.idle":"2025-02-18T05:46:58.431028Z","shell.execute_reply":"2025-02-18T05:46:58.43094Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Predict Inner Function ---\nimport shutil\nimport tempfile\n\ndef predict_inner(problem_statement: str, directory: str) -> Optional[str]:\n\n    # 1. Run file selection logic:\n    directory_string = stringify_directory(directory)\n    selection_completion_texts, file_queries = get_selection_query(directory_string, problem_statement) #Unnecessary for current MCTS implementation.\n    # For simplicity, let's just use the first file_query result:\n    file_query = file_queries[0] if file_queries else {}\n\n    #2. Get only the selected files content to pass into MCTS:\n    selected_files_content = fetch_file_contents(file_query)\n\n    # 3. Create initial repo state as a very long string.\n    file_content_repo: str = \"\"\n    for root, dirs, files in os.walk(directory):\n        for file in files:\n            full_path: str = os.path.join(root, file)\n            with open(full_path, \"r\", encoding=\"utf-8\", errors=\"replace\") as f:\n                file_content_repo += f.read() #This entire file content is no longer passed in.\n\n\n    # 4. Create root node for MCTS with the selected content\n    root = MCTSNode(repo_state= file_content_repo,\n                       patch_history=[],\n                       problem_statement=problem_statement,\n                       selected_files_content = selected_files_content) # Pass the selected content\n\n\n    # 5. Run the MCTS\n    best_patch = mcts(root=root, repo_path=directory, iterations=5) # Run MCTS\n\n    return best_patch","metadata":{"_cell_guid":"41d164a7-440a-4b3d-b82b-9473575c89a4","_uuid":"d5aa7692-14fa-46cb-81aa-4a05d2990df4","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.432431Z","iopub.status.idle":"2025-02-18T05:46:58.432647Z","shell.execute_reply":"2025-02-18T05:46:58.432561Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Predict Function (Outer Wrapper) ---\n\ndef predict(problem_statement: str, repo_archive: io.BytesIO, pip_packages_archive: io.BytesIO, env_setup_cmds_templates: List[str]) -> Optional[str]:\n    \"\"\"Main predict function for inference server.\"\"\"\n    with open(\"repo_archive.tar\", \"wb\") as f:\n        f.write(repo_archive.read())\n\n    repo_path: str = REPO_PATH\n    if os.path.exists(repo_path):\n        shutil.rmtree(repo_path)\n    shutil.unpack_archive(\"repo_archive.tar\", extract_dir=repo_path)\n    os.remove(\"repo_archive.tar\")\n\n    patch_string: Optional[str] = None\n    patch_string = predict_inner(\n        problem_statement=problem_statement, directory=repo_path\n    )\n    shutil.rmtree(repo_path)\n    print(\"submitted patch_string\")\n    print(patch_string)\n    return patch_string","metadata":{"_cell_guid":"f6603ce1-36ca-4c32-b6b7-55d970a38851","_uuid":"1fe33c0a-f4bb-4d58-a31a-8a35efed47f4","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.433137Z","iopub.status.idle":"2025-02-18T05:46:58.433393Z","shell.execute_reply":"2025-02-18T05:46:58.433286Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Optional: Data Loading and Demo (for Interactive Testing) ---\nimport os\nimport zipfile\n\nos.makedirs(\"/kaggle/tmp/konwinski-prize-alt\", exist_ok=True)\n\ntry:\n    with zipfile.ZipFile(\"/kaggle/input/konwinski-prize/data.a_zip\", \"r\") as zip_ref:\n        zip_ref.extractall(\"/kaggle/tmp/konwinski-prize-alt/\")\nexcept:\n    pass\n\nimport pandas as pd\n\ndef get_problem(problem_index: int) -> Tuple[str, str, io.BytesIO]:\n    df = pd.read_parquet(\"/kaggle/tmp/konwinski-prize-alt/data/data.parquet\")\n\n    problem_statement: str = df[\"problem_statement\"][problem_index]\n    repo_path: str = f\"/kaggle/tmp/konwinski-prize-alt/data/repos/repo__{df['instance_id'][problem_index]}\"\n\n    import shutil\n    import tempfile\n\n    with tempfile.TemporaryDirectory() as tmpdir:\n        shutil.make_archive(os.path.join(tmpdir, \"a_repo\"), \"tar\", repo_path)\n        with open(os.path.join(tmpdir, \"a_repo.tar\"), \"rb\") as f:\n            repo_archive = io.BytesIO(f.read())\n\n    return problem_statement, repo_path, repo_archive\n\ndemo_problem_index: int = 0\n\nif os.getenv(\"KAGGLE_KERNEL_RUN_TYPE\") == \"Interactive\" and not os.getenv(\n    \"KAGGLE_IS_COMPETITION_RERUN\"\n):\n    problem_statement, repo_path, repo_archive = get_problem(\n        problem_index=demo_problem_index\n    )\n\n    print(repo_path)\n    print(problem_statement)\n    print(len(list(repo_archive)))\n    print(len(list(repo_archive)))","metadata":{"_cell_guid":"91157abf-a02f-47f2-b6b2-64057f65adaa","_uuid":"6d66b2c6-91e0-4652-b671-a2534be7e103","collapsed":false,"execution":{"iopub.status.busy":"2025-02-18T05:46:58.433817Z","iopub.status.idle":"2025-02-18T05:46:58.434023Z","shell.execute_reply":"2025-02-18T05:46:58.433941Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"skip_prediction = False","metadata":{"_cell_guid":"697c6852-96c5-4c3f-85af-64b340b1c65c","_uuid":"6a96bb2f-85cc-4b37-b044-1ecbf8422efd","collapsed":false,"execution":{"iopub.execute_input":"2025-02-18T05:32:19.228846Z","iopub.status.busy":"2025-02-18T05:32:19.22865Z","iopub.status.idle":"2025-02-18T05:32:19.239883Z","shell.execute_reply":"2025-02-18T05:32:19.239174Z","shell.execute_reply.started":"2025-02-18T05:32:19.22883Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Optional: Run Prediction and Evaluate Locally ---\nif os.getenv(\"KAGGLE_KERNEL_RUN_TYPE\") == \"Interactive\" and not os.getenv(\n    \"KAGGLE_IS_COMPETITION_RERUN\"\n):\n    skip_prediction = False\n    problem_statement, repo_path, repo_archive = get_problem(\n        problem_index=demo_problem_index\n    )\n    patch_string = predict(problem_statement, repo_archive, io.BytesIO(), [])\n\nif (\n    os.getenv(\"KAGGLE_KERNEL_RUN_TYPE\") == \"Interactive\"\n    and not os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\")\n    and patch_string is not None\n):\n    import polars as pl\n    df = pl.read_parquet(\"/kaggle/tmp/konwinski-prize-alt/data/data.parquet\")\n    import kaggle_evaluation.konwinski_prize_gateway\n\n    k_prize_gateway = kaggle_evaluation.konwinski_prize_gateway.KPrizeGateway()\n    k_prize_gateway.unpack_data_paths()\n\n    results = k_prize_gateway._evaluate_instance(\n        instance=df.row(demo_problem_index, named=True),\n        patch=patch_string,\n    )\n\n    from collections import Counter\n    print(\n        demo_problem_index, Counter(result.unit_test_outcome for result in results[1:])\n    )\n\nif (\n    os.getenv(\"KAGGLE_KERNEL_RUN_TYPE\") == \"Interactive\"\n    and not os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\") and patch_string is not None\n):\n    from kaggle_evaluation.konwinski_prize_gateway import UnitTestOutcome\n\n    print(\"\\n--- Unit Test Results (if any failed) ---\")\n    for result in results[1:]: # Skip the first result (instance setup)\n        if result.unit_test_outcome != UnitTestOutcome.PASSED:\n            print(f\"Test Name: {result.test_name}\")\n            print(f\"Fail Description:\\n{result.fail_description}\\n\")","metadata":{"_cell_guid":"bbbf2217-ac29-4d86-bffd-afe648ae1167","_kg_hide-output":true,"_uuid":"e03ade0d-f48f-4682-888e-bd8d258767f3","collapsed":false,"execution":{"iopub.execute_input":"2025-02-18T05:32:19.24075Z","iopub.status.busy":"2025-02-18T05:32:19.24055Z","iopub.status.idle":"2025-02-18T05:44:23.857904Z","shell.execute_reply":"2025-02-18T05:44:23.856944Z","shell.execute_reply.started":"2025-02-18T05:32:19.240733Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"skip_prediction = False\ninference_server = (\n    kaggle_evaluation.konwinski_prize_inference_server.KPrizeInferenceServer(\n        get_number_of_instances, predict\n    )\n)\n\nif os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"):\n    inference_server.serve()\nelse:\n    inference_server.run_local_gateway(\n        data_paths=(\n            \"/kaggle/input/konwinski-prize/\",  # Path to the entire competition dataset\n            \"/kaggle/tmp/konwinski-prize/\",  # Path to a scratch directory for unpacking data.a_zip.\n        )  # type: ignore\n    )","metadata":{"_cell_guid":"b73f8a85-f747-4a0d-9598-6bf418752bc6","_uuid":"86683ecc-947a-444b-b073-352bf009e91c","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"outputs":[],"execution_count":null}]}