{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59575,"databundleVersionId":8060720,"sourceType":"competition"},{"sourceId":8479599,"sourceType":"datasetVersion","datasetId":4517815},{"sourceId":8570631,"sourceType":"datasetVersion","datasetId":5124645},{"sourceId":8570635,"sourceType":"datasetVersion","datasetId":5124648},{"sourceId":8582681,"sourceType":"datasetVersion","datasetId":5124768},{"sourceId":8602900,"sourceType":"datasetVersion","datasetId":5147392},{"sourceId":8602901,"sourceType":"datasetVersion","datasetId":5147393},{"sourceId":8602903,"sourceType":"datasetVersion","datasetId":5147395},{"sourceId":8602905,"sourceType":"datasetVersion","datasetId":5147397},{"sourceId":8602906,"sourceType":"datasetVersion","datasetId":5147398},{"sourceId":8602911,"sourceType":"datasetVersion","datasetId":5147402},{"sourceId":8602914,"sourceType":"datasetVersion","datasetId":5147404},{"sourceId":8602915,"sourceType":"datasetVersion","datasetId":5147405},{"sourceId":8602917,"sourceType":"datasetVersion","datasetId":5147406},{"sourceId":8602918,"sourceType":"datasetVersion","datasetId":5147407},{"sourceId":8948251,"sourceType":"datasetVersion","datasetId":5132898},{"sourceId":8948269,"sourceType":"datasetVersion","datasetId":5384814},{"sourceId":9047275,"sourceType":"datasetVersion","datasetId":5454784},{"sourceId":181050169,"sourceType":"kernelVersion"},{"sourceId":181471735,"sourceType":"kernelVersion"},{"sourceId":181471775,"sourceType":"kernelVersion"},{"sourceId":181471815,"sourceType":"kernelVersion"},{"sourceId":181476022,"sourceType":"kernelVersion"},{"sourceId":181478062,"sourceType":"kernelVersion"},{"sourceId":181495884,"sourceType":"kernelVersion"},{"sourceId":181495892,"sourceType":"kernelVersion"},{"sourceId":181495899,"sourceType":"kernelVersion"},{"sourceId":181526336,"sourceType":"kernelVersion"},{"sourceId":181526370,"sourceType":"kernelVersion"}],"dockerImageVersionId":30698,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Setup & Config","metadata":{}},{"cell_type":"code","source":"!rm -r /kaggle/working/*\n%cd /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2024-07-27T16:40:27.970952Z","iopub.execute_input":"2024-07-27T16:40:27.971487Z","iopub.status.idle":"2024-07-27T16:40:29.167421Z","shell.execute_reply.started":"2024-07-27T16:40:27.971430Z","shell.execute_reply":"2024-07-27T16:40:29.166042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Kaggle Environment","metadata":{}},{"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/uspto-patent2cpc-dataset\",\n        \"/kaggle/input/uspto-rare-tokens-dataset\",\n        ] \n        + [f\"/kaggle/input/uspto-tokenized-db-dataset-{i}/tokenized-db-{i}\" for i in range(10)]\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":{"execution":{"iopub.status.busy":"2024-07-27T16:40:29.170257Z","iopub.execute_input":"2024-07-27T16:40:29.170734Z","iopub.status.idle":"2024-07-27T16:41:21.837362Z","shell.execute_reply.started":"2024-07-27T16:40:29.170689Z","shell.execute_reply":"2024-07-27T16:41:21.835778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load Library","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2024-07-27T16:41:21.839192Z","iopub.execute_input":"2024-07-27T16:41:21.839582Z","iopub.status.idle":"2024-07-27T16:41:21.848275Z","shell.execute_reply.started":"2024-07-27T16:41:21.839545Z","shell.execute_reply":"2024-07-27T16:41:21.847000Z"},"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 networkx as nx\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 scipy.optimize import linear_sum_assignment\nfrom tqdm import tqdm\n\nimport whoosh_utils\nfrom const import INF, KEY2QUERY, NUM_CPU, QUERY2KEY\nfrom db import CompleteDB, SingleTokenDB, TokinezedDB\nfrom solver import HitBlock, SimulatedAnnealing, State\nfrom utils import compute_ap, evaluate, load_list_bz2, save_list_bz2\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":{"execution":{"iopub.status.busy":"2024-07-27T16:41:21.851638Z","iopub.execute_input":"2024-07-27T16:41:21.852032Z","iopub.status.idle":"2024-07-27T16:42:18.884286Z","shell.execute_reply.started":"2024-07-27T16:41:21.852001Z","shell.execute_reply":"2024-07-27T16:42:18.882952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Preparation","metadata":{}},{"cell_type":"code","source":"if 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\n    # database\n    PATENT2RARE_TOKENS_PATH = \"/kaggle/tmp/uspto-rare-tokens-dataset/db\"\n\n    PATENT2CPC_PATH = \"/kaggle/tmp/uspto-patent2cpc-dataset/db\"\n    \n    TOKENIZED_SPLIT = 10\n    TOKENIZED_DB_PATHES = [f\"/kaggle/tmp/tokenized-db-{i}/db\" for i in range(TOKENIZED_SPLIT)]\n    TOKENIZED_INDEX_PATHES = [\n        f\"/kaggle/input/uspto-tokenized-index-{i}/index.lz4\" for i in range(TOKENIZED_SPLIT)\n    ]\n\n    TRAIN_MODE = \"train\" in TRAIN_PATH\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    TRAIN_INDEX_PATH = \"/kaggle/input/train-index-difficult/index_2500_1M\"\n\n    # database\n    PATENT2RARE_TOKENS_PATH = \"/kaggle/input/rare-tokens/db\"\n\n    PATENT2CPC_PATH = \"/kaggle/input/patent2cpc/db\"\n    \n    TOKENIZED_SPLIT = 10\n    TOKENIZED_DB_PATHES = [\n        f\"/kaggle/input/all-index-per-patent/split/tokenized-db-{i}/db\"\n        for i in range(TOKENIZED_SPLIT)\n    ]\n    TOKENIZED_INDEX_PATHES = [\n        f\"/kaggle/input/all-index-per-patent/split/tokenized-db-{i}/index.lz4\"\n        for i in range(TOKENIZED_SPLIT)\n    ]\n\n    TRAIN_MODE = True","metadata":{"execution":{"iopub.status.busy":"2024-07-27T16:42:18.885917Z","iopub.execute_input":"2024-07-27T16:42:18.886735Z","iopub.status.idle":"2024-07-27T16:42:18.948231Z","shell.execute_reply.started":"2024-07-27T16:42:18.886697Z","shell.execute_reply":"2024-07-27T16:42:18.947251Z"},"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 = set()\nfor i in range(50):\n    all_patents.update(train[f\"target_{i}\"].to_list())\nall_patents.update(train[\"publication_number\"].to_list())\nprint(len(all_patents))\n\n# if TRAIN_MODE:\n#     train = train.head(300)\ntrain.head(1)","metadata":{"execution":{"iopub.status.busy":"2024-07-27T16:42:18.949520Z","iopub.execute_input":"2024-07-27T16:42:18.949997Z","iopub.status.idle":"2024-07-27T16:42:19.123881Z","shell.execute_reply.started":"2024-07-27T16:42:18.949967Z","shell.execute_reply":"2024-07-27T16:42:19.122706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import whoosh\n\nimport abc\nimport collections\nimport json\nimport os\nimport re\nimport subprocess\n\nfrom tqdm import tqdm\n\nimport bz2\nimport json\n\nimport whoosh.analysis\nimport whoosh.collectors\nimport whoosh.fields\nimport whoosh.index\nimport whoosh.matching\nimport whoosh.qparser\nimport whoosh.query\nimport whoosh.scoring\nimport whoosh.util.text\n\n\nNUMBER_REGEX = re.compile(r\"^(\\d+|\\d{1,3}(,\\d{3})*)(\\.\\d+)?$\")\n\n\nclass NumberFilter(whoosh.analysis.Filter):\n    def __call__(self, tokens):\n        for t in tokens:\n            if not NUMBER_REGEX.match(t.text):\n                yield t\n\n\n\ndef _define_uspto_whoosh_schema():\n    BRS_STOPWORDS = [\n        \"an\",\n        \"are\",\n        \"by\",\n        \"for\",\n        \"if\",\n        \"into\",\n        \"is\",\n        \"no\",\n        \"not\",\n        \"of\",\n        \"on\",\n        \"such\",\n        \"that\",\n        \"the\",\n        \"their\",\n        \"then\",\n        \"there\",\n        \"these\",\n        \"they\",\n        \"this\",\n        \"to\",\n        \"was\",\n        \"will\",\n    ]\n\n    # Prevent both stopwords and numbers from ever being indexed.\n    custom_analyzer = whoosh.analysis.StandardAnalyzer(stoplist=BRS_STOPWORDS) | NumberFilter()\n    # Schema fields match a subset of the USPTO list of searchable indexes:\n    # https://ppubs.uspto.gov/pubwebapp/static/pages/searchable-indexes.html\n    schema = whoosh.fields.Schema(\n        id=whoosh.fields.ID(stored=True),\n        ti=whoosh.fields.TEXT(analyzer=custom_analyzer, stored=False),\n        ab=whoosh.fields.TEXT(analyzer=custom_analyzer, stored=False),\n        clm=whoosh.fields.TEXT(analyzer=custom_analyzer, stored=False),\n        detd=whoosh.fields.TEXT(analyzer=custom_analyzer, stored=False),\n        cpc=whoosh.fields.KEYWORD(stored=False, scorable=True),\n    )\n    return schema\n\n\ndef create_index(output_dir, documents, limitmb=5_000):\n    ix = whoosh.index.create_in(output_dir, schema=_define_uspto_whoosh_schema())\n    writer = ix.writer(proces=os.cpu_count(), multisegment=True, limitmb=limitmb)\n\n    for document in tqdm(documents):\n        document = load_list_bz2(document)\n\n        for doc in document:\n            writer.add_document(\n                id=doc[\"publication_number\"],\n                ti=doc[\"title\"],\n                ab=doc[\"abstract\"],\n                clm=doc[\"claims\"],\n                detd=doc[\"description\"],\n                cpc=doc[\"cpc\"],\n            )\n    writer.commit(optimize=True)\n    ix.close()\n    \nmeta = pl.read_parquet(\"/kaggle/input/uspto-explainable-ai/patent_metadata.parquet\")\nmeta = meta.drop_nulls(subset=[\"publication_date\"])\nmeta = meta.filter(meta[\"publication_number\"].is_in(all_patents))\nmeta = meta.with_columns(\n    pl.col(\"publication_date\").dt.year().alias(\"year\"),\n    pl.col(\"publication_date\").dt.month().alias(\"month\"),\n)\nprint(meta.shape)\n\nimport gc\nimport multiprocessing\nimport os\n\nfrom tqdm import tqdm\n\nos.makedirs(\"/kaggle/working/tmp\", exist_ok=True)\n\n# with multiprocessing.Pool(5) as p:\n#     _ = list(tqdm(p.imap(process_patent, meta.group_by([\"year\", \"month\"]))))\nfor (year, month), meta_df in tqdm(meta.group_by([\"year\", \"month\"])):\n    documents = []\n    this_patents_numbers = meta_df[\"publication_number\"].to_list()\n    patents_df = pl.read_parquet(\n        f\"/kaggle/input/uspto-explainable-ai/patent_data/{year}_{month}.parquet\"\n    )\n    common_patents = set(patents_df[\"publication_number\"].to_list()) & set(this_patents_numbers)\n\n    patents_df = patents_df.filter(pl.col(\"publication_number\").is_in(common_patents))\n    meta_df = meta_df.filter(pl.col(\"publication_number\").is_in(common_patents))\n    patents_df = patents_df.sort(\"publication_number\")\n    meta_df = meta_df.sort(\"publication_number\")\n    assert len(patents_df) == len(meta_df)\n\n    for pub, ti, abs, cl, des, cpc in zip(\n        meta_df[\"publication_number\"],\n        patents_df[\"title\"],\n        patents_df[\"abstract\"],\n        patents_df[\"claims\"],\n        patents_df[\"description\"],\n        meta_df[\"cpc_codes\"].to_list(),\n    ):\n        doc = {\n            \"publication_number\": pub,\n            \"title\": ti,\n            \"abstract\": abs,\n            \"claims\": cl,\n            \"description\": des,\n            \"cpc\": cpc,\n        }\n        documents.append(doc)\n\n    save_list_bz2(documents, f\"/kaggle/working/tmp/{year}_{month}.json.bz2\")\n\n    del documents\n    del patents_df\n    del meta_df\n    gc.collect()\n    \nfrom glob import glob\n\ndocuments = glob(\"/kaggle/working/tmp/*.json.bz2\")\noutput_dir = \"test_index\"\nos.makedirs(output_dir, exist_ok=True)\ncreate_index(output_dir=output_dir, documents=documents)","metadata":{"execution":{"iopub.status.busy":"2024-07-27T16:42:19.125839Z","iopub.execute_input":"2024-07-27T16:42:19.126313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_idx = whoosh_utils.load_index(\"test_index\")\nsearcher = whoosh_utils.get_searcher(train_idx)\nqp = whoosh_utils.get_query_parser()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patent2rare_tokens = plyvel.DB(PATENT2RARE_TOKENS_PATH, create_if_missing=False)\npatent2cpc = plyvel.DB(PATENT2CPC_PATH, create_if_missing=False)\ntokenized_db = TokinezedDB(TOKENIZED_DB_PATHES, TOKENIZED_INDEX_PATHES)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_patent_tokens(target: str):\n    data = tokenized_db.get(target)\n    if data is None:\n        return None\n    return data","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if KAGGLE_ENV:\n    with open(\"/kaggle/input/uspto-json/token_counts.json\", \"r\") as f:\n        token_counts = json.load(f)\n    with open(\"/kaggle/input/uspto-json/cpc2count.json\", \"r\") as f:\n        cpc2count = json.load(f)\nelse:\n    with open(\"/kaggle/input/token-counts/token_counts.json\", \"r\") as f:\n        token_counts = json.load(f)\n    with open(\"/kaggle/input/cpc-counts/cpc2count.json\", \"r\") as f:\n        cpc2count = json.load(f)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---\n# Main Part","metadata":{}},{"cell_type":"code","source":"MAX_TOKEN = 25\nMAX_SCORE = 350\n\n\ndef process(target_ids) -> str:\n    target_ids = target_ids[1:]\n    assert len(target_ids) == 50\n\n    patent2tokens = {}\n    for target in target_ids:\n        data = load_patent_tokens(target)\n        if data is None:\n            patent2tokens[target] = {}\n        else:\n            patent2tokens[target] = data\n\n        cpcs = json.loads(patent2cpc.get(target.encode()).decode())\n        patent2tokens[target][\"cpc\"] = cpcs\n\n    query_mat = [[None for _ in range(50)] for _ in range(50)]\n    score_mat = np.full((50, 50), MAX_SCORE)\n    rare_mat = [[None for _ in range(50)] for _ in range(50)]\n    for i in range(50):\n        for j in range(i + 1, 50):\n            target1 = target_ids[i]\n            target2 = target_ids[j]\n            common_tokens = []\n            for key in patent2tokens[target1].keys():\n                if key not in patent2tokens[target2]:\n                    continue\n                tokens1 = set(patent2tokens[target1][key])\n                tokens2 = set(patent2tokens[target2][key])\n                common = list(tokens1 & tokens2)\n                if key == \"cpc\":\n                    common = [f\"cpc:{token}\" for token in common]\n                else:\n                    common = [f\"{KEY2QUERY[key]}:{token}\" for token in common]\n                common_tokens += common\n\n            rare_tokens = []\n            for token in common_tokens:\n                _key, _token = token.split(\":\")\n                if _key == \"cpc\":\n                    count = cpc2count.get(_token, 14_000_000)\n                else:\n                    count = token_counts[QUERY2KEY[_key]].get(_token, 14_000_000)\n                rare_tokens.append((count, _key, _token))\n            rare_tokens = sorted(rare_tokens)\n            rare_mat[i][j] = rare_tokens\n            rare_mat[j][i] = rare_tokens\n\n            this_query = \"\"\n            for _, key, token in rare_tokens[:MAX_TOKEN]:\n                this_token = f'{key}:\"{token}\"'\n                if len(this_query) + len(this_token) >= 390:\n                    break\n                this_query += this_token\n            query_mat[i][j] = f\"({this_query})\"\n            query_mat[j][i] = f\"({this_query})\"\n\n            score = 0\n            for t in range(MAX_TOKEN):\n                if len(rare_tokens) <= t:\n                    count = 14_000_000\n                else:\n                    count, _, _ = rare_tokens[t]\n                score += np.log(count)\n            score_mat[i, j] = score\n            score_mat[j, i] = score\n\n            if len(this_query) >= 390 or score > MAX_SCORE:\n                query_mat[i][j] = None\n                query_mat[j][i] = None\n                score_mat[i, j] = MAX_SCORE\n                score_mat[j, i] = MAX_SCORE\n\n    # マッチング\n    G = nx.Graph()\n    for i in range(50):\n        for j in range(i + 1, 50):\n            G.add_edge(i, j, weight=score_mat[i, j])\n    matching = nx.algorithms.matching.min_weight_matching(G)\n\n    used = set()\n    this_queries = []\n    scores = []\n    pairs = []\n    for i, j in list(matching):\n        if i in used or j in used:\n            continue\n        if query_mat[i][j] is None:\n            continue\n        results = whoosh_utils.execute_query(query_mat[i][j], qp, searcher)\n        results = set(results) & all_patents\n        if len(results) > 2:\n            continue\n\n        scores.append(score_mat[i, j])\n        this_queries.append(query_mat[i][j])\n        used.add(i)\n        used.add(j)\n        pairs.append((i, j))\n\n    single_queries = []\n    while len(single_queries) + len(this_queries) < 25 and len(used) < 50:\n        i = random.choice(list(set(range(50)) - used))\n        used.add(i)\n\n        data = patent2rare_tokens.get(target_ids[i].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        tokens = \"\"\n        for _, token in rare_tokens[:MAX_TOKEN]:\n            key, token = token.split(\":\")\n            this = f'{key}:\"{token}\"'\n            if len(this) + len(tokens) >= 380:\n                break\n            tokens += this\n        this_query = f\"({tokens})\"\n        single_queries.append(this_query)\n\n    can_use_len = 390 * 25 - sum([len(q) for q in single_queries])\n    this_queries = [\"\" for _ in range(len(this_queries))]\n    cur_sum = 0\n    fail_count = 0\n    if len(this_queries) > 0:\n        for i in range(1000000):\n            pair_i, cycle = i % len(this_queries), i // len(this_queries)\n            if cycle >= len(rare_mat[pairs[pair_i][0]][pairs[pair_i][1]]):\n                fail_count += 1\n                if fail_count > 25:\n                    break\n                continue\n            count, key, token = rare_mat[pairs[pair_i][0]][pairs[pair_i][1]][cycle]\n            if cycle >= MAX_TOKEN and count > 14_000_000 * 0.05:\n                fail_count += 1\n                continue\n            this_token = f'{key}:\"{token}\"'\n            if cur_sum + len(this_token) > can_use_len:\n                fail_count += 1\n                if fail_count > 25:\n                    break\n                continue\n            cur_sum += len(this_token)\n            this_queries[pair_i] += this_token\n            fail_count = 0\n    for i in range(len(this_queries)):\n        this_queries[i] = f\"({this_queries[i]})\"\n    this_queries += single_queries\n\n    query = \" OR \".join(this_queries)\n    if len(query) == 0:\n        query = \"hogefugafooooo\"\n    # print(f\"mean={np.mean(scores)}, min={np.min(scores)}, max={np.max(scores)}\")\n    return query\n\n\nwith multiprocessing.Pool(NUM_CPU) as pool:\n    queries = list(tqdm(pool.imap(process, train.iter_rows()), total=len(train)))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 評価\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    aps = evaluate(all_results, list(train.iter_rows()))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -r /kaggle/working/*","metadata":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lens = []\nfor query in queries:\n    lens.append(len(query))\nplt.hist(lens)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}