{"metadata":{"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59575,"databundleVersionId":8060720,"sourceType":"competition"},{"sourceId":8479599,"sourceType":"datasetVersion","datasetId":4517815},{"sourceId":8570631,"sourceType":"datasetVersion","datasetId":5124645},{"sourceId":8582681,"sourceType":"datasetVersion","datasetId":5124768},{"sourceId":8697173,"sourceType":"datasetVersion","datasetId":5215857},{"sourceId":8770817,"sourceType":"datasetVersion","datasetId":5270696},{"sourceId":8770818,"sourceType":"datasetVersion","datasetId":5270697},{"sourceId":8770820,"sourceType":"datasetVersion","datasetId":5270699},{"sourceId":8770821,"sourceType":"datasetVersion","datasetId":5270700},{"sourceId":8770822,"sourceType":"datasetVersion","datasetId":5270701},{"sourceId":8770823,"sourceType":"datasetVersion","datasetId":5270702},{"sourceId":8770824,"sourceType":"datasetVersion","datasetId":5270703},{"sourceId":8770826,"sourceType":"datasetVersion","datasetId":5270705},{"sourceId":8770829,"sourceType":"datasetVersion","datasetId":5270708},{"sourceId":8770830,"sourceType":"datasetVersion","datasetId":5270709},{"sourceId":8770831,"sourceType":"datasetVersion","datasetId":5270710},{"sourceId":8770832,"sourceType":"datasetVersion","datasetId":5270711},{"sourceId":8770834,"sourceType":"datasetVersion","datasetId":5270713},{"sourceId":8770835,"sourceType":"datasetVersion","datasetId":5270714},{"sourceId":8770836,"sourceType":"datasetVersion","datasetId":5270715},{"sourceId":8875137,"sourceType":"datasetVersion","datasetId":5342363},{"sourceId":8875150,"sourceType":"datasetVersion","datasetId":5342374},{"sourceId":8875157,"sourceType":"datasetVersion","datasetId":5342376},{"sourceId":8875159,"sourceType":"datasetVersion","datasetId":5342377},{"sourceId":8875160,"sourceType":"datasetVersion","datasetId":5342378},{"sourceId":8875583,"sourceType":"datasetVersion","datasetId":5342193},{"sourceId":8875743,"sourceType":"datasetVersion","datasetId":5342196},{"sourceId":8875745,"sourceType":"datasetVersion","datasetId":5342194},{"sourceId":8875749,"sourceType":"datasetVersion","datasetId":5342197},{"sourceId":8875750,"sourceType":"datasetVersion","datasetId":5342192},{"sourceId":8948251,"sourceType":"datasetVersion","datasetId":5132898},{"sourceId":9047275,"sourceType":"datasetVersion","datasetId":5454784},{"sourceId":181050169,"sourceType":"kernelVersion"},{"sourceId":185202816,"sourceType":"kernelVersion"},{"sourceId":185203355,"sourceType":"kernelVersion"},{"sourceId":185203360,"sourceType":"kernelVersion"},{"sourceId":185203366,"sourceType":"kernelVersion"},{"sourceId":185203371,"sourceType":"kernelVersion"},{"sourceId":185204828,"sourceType":"kernelVersion"},{"sourceId":185205510,"sourceType":"kernelVersion"},{"sourceId":185205876,"sourceType":"kernelVersion"},{"sourceId":185206618,"sourceType":"kernelVersion"},{"sourceId":185207674,"sourceType":"kernelVersion"},{"sourceId":185208544,"sourceType":"kernelVersion"},{"sourceId":185208566,"sourceType":"kernelVersion"},{"sourceId":185209912,"sourceType":"kernelVersion"},{"sourceId":185209929,"sourceType":"kernelVersion"},{"sourceId":185210713,"sourceType":"kernelVersion"}],"dockerImageVersionId":30698,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false},"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.13"},"papermill":{"default_parameters":{},"duration":809.091152,"end_time":"2024-07-06T05:34:30.750181","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-07-06T05:21:01.659029","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Setup & Config","metadata":{"papermill":{"duration":0.012025,"end_time":"2024-07-06T05:21:05.479111","exception":false,"start_time":"2024-07-06T05:21:05.467086","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!rm -r /kaggle/working/*\n%cd /kaggle/working","metadata":{"papermill":{"duration":1.203869,"end_time":"2024-07-06T05:21:06.694555","exception":false,"start_time":"2024-07-06T05:21:05.490686","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T16:56:06.693261Z","iopub.execute_input":"2024-07-27T16:56:06.694433Z","iopub.status.idle":"2024-07-27T16:56:07.887231Z","shell.execute_reply.started":"2024-07-27T16:56:06.694381Z","shell.execute_reply":"2024-07-27T16:56:07.886001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Kaggle Environment","metadata":{"papermill":{"duration":0.010757,"end_time":"2024-07-06T05:21:06.716643","exception":false,"start_time":"2024-07-06T05:21:06.705886","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nfrom tqdm import tqdm\n\nKAGGLE_ENV = not os.path.exists(\"/kaggle/.vscode\")\nENV_NAME = \"kaggle\" if KAGGLE_ENV else \"local\"\nprint(f\"{KAGGLE_ENV=}\")\nprint(f\"{ENV_NAME=}\")\n\nif KAGGLE_ENV:\n    !pip install -U -q plyvel --no-index --find-links=file:///kaggle/input/uspto-gen-wheel/plyvel\n\n    move_dirs = (\n        [\n        \"/kaggle/input/preprocess-all-token-single\",\n        \"/kaggle/input/uspto-rare-tokens-dataset\",\n        ] \n        + [f\"/kaggle/input/complete-db-{i}\" for i in range(15)]\n        + [f\"/kaggle/input/complete-db-v2-{i}\" for i in range(5)]\n        + [f\"/kaggle/input/uspto-ratio-db-{i}\" for i in range(20)]\n    )\n\n    !mkdir /kaggle/tmp\n    for move_dir in tqdm(move_dirs):\n        !cp -r {move_dir} /kaggle/tmp/{move_dir.split(\"/\")[-1]}\n    !ls /kaggle/tmp","metadata":{"papermill":{"duration":625.217332,"end_time":"2024-07-06T05:31:31.945033","exception":false,"start_time":"2024-07-06T05:21:06.727701","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T16:56:07.889948Z","iopub.execute_input":"2024-07-27T16:56:07.890291Z","iopub.status.idle":"2024-07-27T16:58:51.152987Z","shell.execute_reply.started":"2024-07-27T16:56:07.890255Z","shell.execute_reply":"2024-07-27T16:58:51.150876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load Library","metadata":{"papermill":{"duration":0.015424,"end_time":"2024-07-06T05:31:31.977664","exception":false,"start_time":"2024-07-06T05:31:31.962240","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport sys\n\nif KAGGLE_ENV:\n    PACKAGE_DIR = \"/kaggle/input/uspto-src-for-public/src\"\nelse:\n    PACKAGE_DIR = \"/kaggle/src\"\nsys.path.append(PACKAGE_DIR)\nsys.path.append(os.path.join(PACKAGE_DIR, \"Penguin-ML-Library\"))","metadata":{"papermill":{"duration":0.034053,"end_time":"2024-07-06T05:31:32.028987","exception":false,"start_time":"2024-07-06T05:31:31.994934","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T16:58:51.156804Z","iopub.execute_input":"2024-07-27T16:58:51.157446Z","iopub.status.idle":"2024-07-27T16:58:51.169609Z","shell.execute_reply.started":"2024-07-27T16:58:51.157375Z","shell.execute_reply":"2024-07-27T16:58:51.168251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nimport multiprocessing\nimport random\nimport warnings\nfrom collections import Counter\nfrom typing import List, Set, Tuple\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport plyvel\nimport polars as pl\nimport yaml\nfrom penguinml.utils.logger import get_logger, init_logger\nfrom penguinml.utils.set_seed import seed_base\nfrom penguinml.utils.timer import Timer\nfrom tqdm import tqdm\n\nimport whoosh_utils\nfrom const import INF, NUM_CPU\nfrom db import CompleteDB, SingleTokenDB\nfrom solver import HitBlock, SimulatedAnnealing, State\nfrom utils import compute_ap, evaluate\n\nwarnings.filterwarnings(\"ignore\")\nMODEL_NAME = \"baseline\"\nCFG = yaml.safe_load(open(os.path.join(PACKAGE_DIR, \"config.yaml\"), \"r\"))\nprint(CFG[MODEL_NAME][\"execution\"][\"exp_id\"])\nCFG[\"output_dir\"] = f\"/kaggle/output/{CFG[MODEL_NAME]['execution']['exp_id']}\"\n# !rm -r {CFG[\"output_dir\"]}\nos.makedirs(CFG[\"output_dir\"], exist_ok=True)\n\ninit_logger(\"log.log\")\nlogger = get_logger(\"main\")\nseed_base(CFG[MODEL_NAME][\"execution\"][\"seed\"])","metadata":{"papermill":{"duration":62.308476,"end_time":"2024-07-06T05:32:34.355774","exception":false,"start_time":"2024-07-06T05:31:32.047298","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T16:58:51.171492Z","iopub.execute_input":"2024-07-27T16:58:51.171949Z","iopub.status.idle":"2024-07-27T16:59:47.350882Z","shell.execute_reply.started":"2024-07-27T16:58:51.171910Z","shell.execute_reply":"2024-07-27T16:59:47.349584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Preparation","metadata":{"papermill":{"duration":0.018261,"end_time":"2024-07-06T05:32:34.390857","exception":false,"start_time":"2024-07-06T05:32:34.372596","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# DANGER_TYPE = \"all\"\nDANGER_TYPE = \"2hop\"\nassert DANGER_TYPE in [\"all\", \"2hop\"]\n\nif KAGGLE_ENV:\n    TRAIN_PATH = \"/kaggle/input/uspto-explainable-ai/test.csv\"\n    if len(pl.read_csv(TRAIN_PATH)) < 100:\n        TRAIN_PATH = \"/kaggle/input/uspto-train-index-2500/train2500_seed0.parquet\"\n    TRAIN_INDEX_PATH = \"/kaggle/input/uspto-train-index-2500/index_2500_200k\"\n    NN_DF_PATH = \"/kaggle/input/uspto-explainable-ai/nearest_neighbors.csv\"\n\n    # database\n    PATENT2RARE_TOKENS_PATH = \"/kaggle/tmp/uspto-rare-tokens-dataset/db\"\n\n    COMPLETE_DB_PATH = [f\"/kaggle/tmp/complete-db-{i}/complete-db-{i}/db\" for i in range(15)]\n    COMPLETE_INDEX_PATH = [f\"/kaggle/input/uspto-complete-index-{i}/index.lz4\" for i in range(15)]\n    \n    COMPLETE_DB_V2_PATH = [\n        f\"/kaggle/tmp/complete-db-v2-{i}/db\" for i in range(5)\n    ]\n    COMPLETE_INDEX_V2_PATH = [\n        f\"/kaggle/input/complete-db-index-v2-{i}/index.lz4\" for i in range(5)\n    ]\n    \n    SINGLE_TOKEN_DB_PATH = \"/kaggle/tmp/preprocess-all-token-single/db/db\"\n    SINGLE_TOKEN_INDEX_PATH = \"/kaggle/tmp/preprocess-all-token-single/index.lz4\"\n\n    TRAIN_MODE = \"train\" in TRAIN_PATH\n    VISUALIZE = False\nelse:\n    TRAIN_PATH = \"/kaggle/input/uspto-train-data-2500/train2500_seed0.parquet\"\n    TRAIN_INDEX_PATH = \"/kaggle/input/train-index-2500/index_2500_200k\"\n    NN_DF_PATH = \"/kaggle/input/uspto-boolean-search-optimization/nearest_neighbors.csv\"\n\n    # database\n    PATENT2RARE_TOKENS_PATH = \"/kaggle/input/rare-tokens/db\"\n\n    COMPLETE_DB_PATH = [\n        f\"/kaggle/input/preprocess-complete/split/complete-db-{i}/db\" for i in range(15)\n    ]\n    COMPLETE_INDEX_PATH = [\n        f\"/kaggle/input/preprocess-complete/split/complete-db-{i}/index.lz4\" for i in range(15)\n    ]\n\n    COMPLETE_DB_V2_PATH = [\n        f\"/kaggle/input/preprocess-complete-v2/split/complete-db-{i}/db\" for i in range(5)\n    ]\n    COMPLETE_INDEX_V2_PATH = [\n        f\"/kaggle/input/preprocess-complete-v2/split/complete-db-{i}/index.lz4\" for i in range(5)\n    ]\n\n    SINGLE_TOKEN_DB_PATH = \"/kaggle/input/preprocess-all-token-single/db\"\n    SINGLE_TOKEN_INDEX_PATH = \"/kaggle/input/preprocess-all-token-single/index.lz4\"\n\n    TRAIN_MODE = True\n    VISUALIZE = False","metadata":{"papermill":{"duration":0.041169,"end_time":"2024-07-06T05:32:34.450544","exception":false,"start_time":"2024-07-06T05:32:34.409375","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T17:27:06.582725Z","iopub.execute_input":"2024-07-27T17:27:06.584958Z","iopub.status.idle":"2024-07-27T17:27:06.614307Z","shell.execute_reply.started":"2024-07-27T17:27:06.584896Z","shell.execute_reply":"2024-07-27T17:27:06.612817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN_PATH.split(\".\")[-1] == \"parquet\":\n    train = pl.read_parquet(TRAIN_PATH)\nelse:\n    train = pl.read_csv(TRAIN_PATH)\n\nif TRAIN_MODE:\n    train = train.filter(~train[\"publication_number\"].str.starts_with(\"US-D\"))\n\nall_patents = []\nfor i in range(50):\n    all_patents += train[f\"target_{i}\"].to_list()\nall_patents = set(all_patents)\n\nif DANGER_TYPE == \"2hop\":\n    all_df = pl.read_csv(NN_DF_PATH)\n    all_df = all_df.filter(all_df[\"publication_number\"].is_in(all_patents))\n    for i in range(50):\n        all_patents.update(all_df[f\"neighbor_{i}\"].to_list())\nprint(len(all_patents))\n\nif TRAIN_MODE:\n    train = train.head(300)\ntrain.head(1)","metadata":{"papermill":{"duration":71.289226,"end_time":"2024-07-06T05:33:45.758904","exception":false,"start_time":"2024-07-06T05:32:34.469678","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T16:59:47.421763Z","iopub.execute_input":"2024-07-27T16:59:47.422158Z","iopub.status.idle":"2024-07-27T17:00:49.936920Z","shell.execute_reply.started":"2024-07-27T16:59:47.422126Z","shell.execute_reply":"2024-07-27T17:00:49.935611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN_MODE:\n    train_idx = whoosh_utils.load_index(TRAIN_INDEX_PATH)\n    searcher = whoosh_utils.get_searcher(train_idx)\n    qp = whoosh_utils.get_query_parser()","metadata":{"papermill":{"duration":0.033858,"end_time":"2024-07-06T05:33:45.809809","exception":false,"start_time":"2024-07-06T05:33:45.775951","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T17:00:49.938427Z","iopub.execute_input":"2024-07-27T17:00:49.938828Z","iopub.status.idle":"2024-07-27T17:01:43.541060Z","shell.execute_reply.started":"2024-07-27T17:00:49.938795Z","shell.execute_reply":"2024-07-27T17:01:43.539513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patent2rare_tokens = plyvel.DB(PATENT2RARE_TOKENS_PATH, create_if_missing=False)\n\ncomplete_db = CompleteDB(COMPLETE_DB_PATH, COMPLETE_INDEX_PATH)\ncomplete_v2_db = CompleteDB(COMPLETE_DB_V2_PATH, COMPLETE_INDEX_V2_PATH)\nsingle_token_db = SingleTokenDB(SINGLE_TOKEN_DB_PATH, SINGLE_TOKEN_INDEX_PATH)","metadata":{"papermill":{"duration":0.993251,"end_time":"2024-07-06T05:33:46.826517","exception":false,"start_time":"2024-07-06T05:33:45.833266","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T17:01:43.542940Z","iopub.execute_input":"2024-07-27T17:01:43.543407Z","iopub.status.idle":"2024-07-27T17:01:44.075809Z","shell.execute_reply.started":"2024-07-27T17:01:43.543364Z","shell.execute_reply":"2024-07-27T17:01:44.074564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---\n# Main Part","metadata":{"papermill":{"duration":0.019062,"end_time":"2024-07-06T05:33:46.864967","exception":false,"start_time":"2024-07-06T05:33:46.845905","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def enumerate_token_queries(\n    target_ids: List[str],\n    center_id: str,\n    top_k: int = 30,\n) -> List[HitBlock]:\n    \"\"\"\n    3. cpcを使わずにねじ込めるqueryを追加\n    \"\"\"\n    cands = []\n    # single patent\n    for target_id in target_ids:\n        data = patent2rare_tokens.get(target_id.encode())\n        if data is None:\n            continue\n        rare_tokens = json.loads(data.decode())\n        if rare_tokens is None or len(rare_tokens) == 0:\n            continue\n\n        if rare_tokens[0][0] == 1:\n            score = -INF\n            tokens = [rare_tokens[0][1]]\n        else:\n            score = 1\n            tokens = []\n            for freq, token in rare_tokens:\n                score *= freq / 10_000_000\n                tokens.append(token)\n                if score < 1e-13:\n                    break\n        this_query = f\"({' '.join(tokens)})\"\n        block = HitBlock(this_query, {target_id}, 1, 0, 0)\n        cands.append((score, block))\n    cands = sorted(cands, key=lambda x: x[0])[:top_k]\n    cands = [cand[1] for cand in cands]\n\n    # multi patents\n    data = single_token_db.get(center_id)\n    if data is not None:\n        for token, n_inner, n_outer, pattents in data:\n            if n_outer < 5:\n                assert n_inner == len(pattents)\n                block = HitBlock(f\"({token})\", set(pattents), n_inner, 0, n_outer)\n                cands.append(block)\n    return cands","metadata":{"papermill":{"duration":0.040101,"end_time":"2024-07-06T05:33:46.923667","exception":false,"start_time":"2024-07-06T05:33:46.883566","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T17:01:44.077577Z","iopub.execute_input":"2024-07-27T17:01:44.078526Z","iopub.status.idle":"2024-07-27T17:01:44.092945Z","shell.execute_reply.started":"2024-07-27T17:01:44.078485Z","shell.execute_reply":"2024-07-27T17:01:44.091689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def optimize(\n    use_cpcs: List[str],\n    query_patents: List[List[HitBlock]],\n    center_id: str,\n    target_ids: List[str],\n) -> Tuple[str, int, float]:\n    \"\"\"最適化\"\"\"\n    if TRAIN_MODE:\n        timer = Timer()\n        log = \"\"\n        log_file = f\"/kaggle/working/log/{center_id}.txt\"\n        os.makedirs(os.path.dirname(log_file), exist_ok=True)\n\n    if TRAIN_MODE and VISUALIZE:\n        visualize_log = {}\n\n    state = State(use_cpcs, query_patents, target_ids)\n\n    global_best_state = state.copy()\n    global_best_score = -INF\n    global_best_query = \"\"\n    global_best_hit_patent_count = 0\n    global_best_ap = 0\n\n    n_iter = 50000\n    sa = SimulatedAnnealing(0.001, 0.00001, n_iter)\n\n    best_score = -INF\n    best_query = \"\"\n    best_hit_patent_count = 0\n    best_ap = 0\n    last_updated = 0\n\n    for now_iter in range(n_iter):\n        if len(query_patents) == 0:\n            continue\n\n        # 近傍に遷移\n        flipped_cpc_queries = []\n        if random.random() < 0.5:\n            # single flip\n            r_cpc = random.randrange(len(query_patents))\n            if len(query_patents[r_cpc]) == 0:\n                continue\n            r_query = np.random.choice(len(query_patents[r_cpc]))\n            state.bit_flip(r_cpc, r_query)\n            flipped_cpc_queries.append((r_cpc, r_query))\n        else:\n            # off -> on (high probability)\n            r_cpc1 = random.randrange(len(query_patents))\n            if len(query_patents[r_cpc1]) == 0:\n                continue\n            r_query1 = np.random.choice(len(query_patents[r_cpc1]))\n\n            # on -> off\n            used_query_indices = list(state.used_query_indices)\n            if len(used_query_indices) == 0:\n                continue\n            r_cpc2, r_query2 = random.choice(used_query_indices)\n\n            state.bit_flip(r_cpc1, r_query1)\n            state.bit_flip(r_cpc2, r_query2)\n            flipped_cpc_queries.append((r_cpc1, r_query1))\n            flipped_cpc_queries.append((r_cpc2, r_query2))\n\n        # 評価\n        score = state.evaluate()\n        d_worsen = best_score - score\n        if sa.accept(d_worsen):\n            if best_score < score:\n                last_updated = now_iter\n                best_score = score\n                best_query = state.get_query()\n                best_hit_patent_count = state.hit_patent_count\n                best_ap = state.ap()\n\n            if best_score > global_best_score:\n                global_best_state = state.copy()\n                global_best_score = best_score\n                global_best_query = best_query\n                global_best_hit_patent_count = best_hit_patent_count\n                global_best_ap = best_ap\n        else:\n            for r_cpc, r_query in flipped_cpc_queries:\n                state.bit_flip(r_cpc, r_query)\n\n        if TRAIN_MODE and VISUALIZE:\n            visualize_log[now_iter] = {\n                \"score\": score,\n                \"used_query_indices\": list(state.used_query_indices),\n            }\n\n        # restart\n        if now_iter - last_updated > 1000:\n            best_score = -INF\n            best_query = \"\"\n            best_hit_patent_count = 0\n            best_ap = 0\n            last_updated = now_iter\n            state = global_best_state.copy()\n\n    # log\n    if TRAIN_MODE:\n        # optimize result\n        log += f\"Query: {global_best_query}\\n\"\n        log += f\"[Optimization]\\n\"\n        log += f\"AP: {global_best_ap}\\n\"\n        log += f\"Hit Patent Count: {global_best_hit_patent_count}\\n\"\n        log += f\"SA Score: {global_best_score}\\n\"\n        time = timer.elapsed_sec()\n        log += f\"Time: {time:.2f} sec\\n\"\n\n        # search result\n        log += f\"[Strict]\\n\"\n        results = whoosh_utils.execute_query(global_best_query, qp, searcher)\n        ap = compute_ap(results, target_ids)\n        log += f\"AP: {ap}\\n\"\n        log += f\"Hit Patent Count: {len(set(results) & set(target_ids))}\\n\"\n        log += f\"Hit Patents: {results}\\n\"\n        corrects = [r in target_ids for r in results]\n        log += f\"Corrects: {corrects}\\n\"\n\n        with open(log_file, \"w\") as f:\n            f.write(log)\n\n    if TRAIN_MODE and VISUALIZE:\n        with open(f\"/kaggle/working/log/{center_id}.json\", \"w\") as f:\n            json.dump(visualize_log, f)\n    return (global_best_query, global_best_hit_patent_count, global_best_ap)","metadata":{"papermill":{"duration":0.050585,"end_time":"2024-07-06T05:33:46.993280","exception":false,"start_time":"2024-07-06T05:33:46.942695","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T17:01:44.095193Z","iopub.execute_input":"2024-07-27T17:01:44.095755Z","iopub.status.idle":"2024-07-27T17:01:44.124373Z","shell.execute_reply.started":"2024-07-27T17:01:44.095691Z","shell.execute_reply":"2024-07-27T17:01:44.123279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Problem:\n    def __init__(self, target_ids: List[str]):\n        # parse target_ids\n        self.center_id = target_ids[0]\n        self.target_ids = list(target_ids)[1:]\n        assert len(self.target_ids) == 50\n\n        # 3. 強力な前処理\n        cpc2idx = {}\n        self.use_cpcs = []\n        self.query_patents = []\n        data = []\n        d = complete_db.get(self.center_id)\n        if d is not None:\n            data += d\n        d = complete_v2_db.get(self.center_id)\n        if d is not None:\n            data += d\n        if data is not None:\n            all_data = []\n            for cpc, token, inner, outer in data:\n                all_data.append((f\"cpc:{cpc}\", token, inner, outer))\n                all_data.append((token, f\"cpc:{cpc}\", inner, outer))  # reversed\n            data = sorted(\n                all_data, key=lambda x: (len(x[3]), -len(x[2]))\n            )  # (n_outer, Reversed(n_innver))\n            used = []\n            for cpc, token, inner, outer in data:\n                key = set(inner)\n                ok = True\n                for used_key in used:\n                    if key.issubset(used_key):\n                        ok = False\n                        break\n                if not ok:\n                    continue\n                used.append(key)\n\n                if DANGER_TYPE == \"2hop\":\n                    n_danger = len(set(outer) & all_patents)\n                    block = HitBlock(\n                        token=token,\n                        inner_patents=set(inner),\n                        n_inner=len(inner),\n                        n_outer=len(outer),\n                        n_danger=n_danger,\n                    )\n                elif DANGER_TYPE == \"all\":\n                    block = HitBlock(\n                        token=token,\n                        inner_patents=set(inner),\n                        n_inner=len(inner),\n                        n_outer=0,\n                        n_danger=len(outer),\n                    )\n                else:\n                    raise ValueError\n\n                if cpc not in cpc2idx:\n                    cpc2idx[cpc] = len(cpc2idx)\n                    self.use_cpcs.append(cpc)\n                    self.query_patents.append([])\n                self.query_patents[cpc2idx[cpc]].append(block)\n        # _sum = sum(len(qp) for qp in self.query_patents)\n        # print(f\"{_sum=}\")\n\n        # 4. cpcを使わずにねじ込めるqueryを追加\n        self.use_cpcs.append(None)\n        self.query_patents.append(enumerate_token_queries(self.target_ids, self.center_id))\n\n    def solve(self):\n        best_query, best_hit_patent_count, best_ap = optimize(\n            self.use_cpcs,\n            self.query_patents,\n            self.center_id,\n            self.target_ids,\n        )\n        return best_query, best_hit_patent_count, best_ap\n\n\ndef _process(problem: Problem):\n    return problem.solve()\n\n\ndef solve_problems(problems: List[Problem]):\n    with multiprocessing.Pool(NUM_CPU) as pool:\n        results = list(\n            tqdm(\n                pool.imap(_process, problems),\n                total=len(problems),\n                desc=\"Solve Problems\",\n            )\n        )\n\n    queries = []\n    counts = []\n    aps = []\n    for q, c, ap in results:\n        queries.append(q)\n        counts.append(c)\n        aps.append(ap)\n    return queries, counts, aps","metadata":{"papermill":{"duration":0.045315,"end_time":"2024-07-06T05:33:47.055514","exception":false,"start_time":"2024-07-06T05:33:47.010199","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T17:01:44.126106Z","iopub.execute_input":"2024-07-27T17:01:44.126526Z","iopub.status.idle":"2024-07-27T17:01:44.338846Z","shell.execute_reply.started":"2024-07-27T17:01:44.126485Z","shell.execute_reply":"2024-07-27T17:01:44.337402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 最適化\nproblems = []\nfor target_ids in tqdm(train.iter_rows(), desc=\"Generate Problems\"):\n    problem = Problem(target_ids)\n    problems.append(problem)\n\nqueries, counts, aps = solve_problems(problems)\nprint(\"== SA Result ==\")\nprint(\"AP:\", np.mean(aps))\nprint(\"Count:\", np.mean(counts))\n\n# 評価\nif TRAIN_MODE:\n    print(\"== Evaluation ==\")\n    all_results = []\n    for query in tqdm(queries):\n        results = whoosh_utils.execute_query(query, qp, searcher)\n        all_results.append(results)\n    evaluate(all_results, list(train.iter_rows()))","metadata":{"papermill":{"duration":40.834972,"end_time":"2024-07-06T05:34:27.910141","exception":false,"start_time":"2024-07-06T05:33:47.075169","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-27T17:27:09.263486Z","iopub.execute_input":"2024-07-27T17:27:09.263995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# submit\ntrain = train.with_columns(pl.Series(\"query\", queries)).select([\"publication_number\", \"query\"])\ntrain.write_csv(\"submission.csv\")\ntrain.head(1)","metadata":{"papermill":{"duration":0.054658,"end_time":"2024-07-06T05:34:27.992828","exception":false,"start_time":"2024-07-06T05:34:27.938170","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}