{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":84795,"databundleVersionId":10462807,"sourceType":"competition"},{"sourceId":211459,"sourceType":"modelInstanceVersion","modelInstanceId":180283,"modelId":202544}],"dockerImageVersionId":30822,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np \nimport io\nimport os\nimport shutil\n\nimport pandas as pd\nimport polars as pl\n\nimport kaggle_evaluation.konwinski_prize_inference_server","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-27T12:59:02.769375Z","iopub.execute_input":"2024-12-27T12:59:02.769746Z","iopub.status.idle":"2024-12-27T12:59:05.579849Z","shell.execute_reply.started":"2024-12-27T12:59:02.769718Z","shell.execute_reply":"2024-12-27T12:59:05.577845Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"instance_count = None\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":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-27T12:59:05.580703Z","iopub.status.idle":"2024-12-27T12:59:05.581116Z","shell.execute_reply":"2024-12-27T12:59:05.580968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\n# 假设 repo_path 是解压后的仓库路径\n# repo_path = 'repo'\ndef generate_source_code(repo_path):\n    # 用来存储所有文件的内容\n    source_code = \"\"\n\n    # 遍历仓库中的所有 .py 文件\n    for root, dirs, files in os.walk(repo_path):\n        for file in files:\n            if file.endswith(\".py\"):  # 只选择 Python 文件\n                file_path = os.path.join(root, file)\n\n                # 读取文件内容\n                with open(file_path, 'r', encoding='utf-8') as f:\n                    file_content = f.read()\n                    source_code += f\"\\n\\n# Start of file: {file_path}\\n\"  # 为每个文件加上标识\n                    source_code += file_content  # 将文件内容添加到 source_code 中\n\n    # 打印部分文件内容进行检查\n    # print(source_code[:2000])  # 打印前 2000 字符，避免输出过长\n    return source_code","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-27T12:59:05.582436Z","iopub.status.idle":"2024-12-27T12:59:05.582782Z","shell.execute_reply":"2024-12-27T12:59:05.582652Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import io\nimport os\nimport shutil\nimport subprocess\nimport torch\nfrom transformers import T5ForConditionalGeneration, RobertaTokenizer\n\ninstance_count = None\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\n\nfirst_prediction = True\n\ndef predict(problem_statement: str, repo_archive: io.BytesIO) -> str:\n    \"\"\"This function will predict the fixed code for the given problem statement.\"\"\"\n    global first_prediction\n    if not first_prediction:\n        return None  # Skip issue.\n\n    # Unpack the repository archive\n    with open('repo_archive.tar', 'wb') as f:\n        f.write(repo_archive.read())\n    repo_path = 'repo'\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    first_prediction = False\n\n    # 指定 tokenizer 文件夹路径\n    tokenizer = RobertaTokenizer.from_pretrained(\"/kaggle/input/mytest/transformers/default1/1\", local_files_only=True)\n    \n    # 加载 safetensors 格式的模型\n    model = T5ForConditionalGeneration.from_pretrained(\"/kaggle/input/mytest/transformers/default1/1\")\n    \n    # 合并 problem_statement 和 source_code\n    input_text = problem_statement +'\\n'+ generate_source_code(repo_path)\n    \n    inputs = tokenizer(input_text, return_tensors=\"pt\", padding=True, truncation=True)\n    \n    # 使用模型生成修复代码\n    generated_code = model.generate(input_ids=inputs['input_ids'], max_length=1024,num_beams=20,\n                                    no_repeat_ngram_size=12,top_p=0.9,temperature=0.4,do_sample=True)\n    \n    # 解码生成的代码\n    generated_code_str = tokenizer.decode(generated_code[0], skip_special_tokens=True)\n\n    # 创建一个 DataFrame\n    submission_df = pd.DataFrame({\n        'id': 1,        # 样本 ID 列\n        'prediction': generated_code_str # 模型的预测结果列\n    })\n\n    # 将 DataFrame 保存为 CSV 文件\n    submission_df.to_csv('/kaggle/working/submission.csv', index=False)\n\n    return generated_code_str","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-27T12:59:05.583739Z","iopub.status.idle":"2024-12-27T12:59:05.584099Z","shell.execute_reply":"2024-12-27T12:59:05.583962Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inference_server = kaggle_evaluation.konwinski_prize_inference_server.KPrizeInferenceServer(\n    get_number_of_instances,   \n    predict\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        )\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-27T12:59:05.584984Z","iopub.status.idle":"2024-12-27T12:59:05.58574Z","shell.execute_reply":"2024-12-27T12:59:05.585442Z"}},"outputs":[],"execution_count":null}]}