{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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":8246447,"sourceType":"datasetVersion","datasetId":4892374},{"sourceId":8323913,"sourceType":"datasetVersion","datasetId":4944579},{"sourceId":8479599,"sourceType":"datasetVersion","datasetId":4517815},{"sourceId":8967222,"sourceType":"datasetVersion","datasetId":5398083},{"sourceId":174185912,"sourceType":"kernelVersion"},{"sourceId":187982647,"sourceType":"kernelVersion"},{"sourceId":189158133,"sourceType":"kernelVersion"}],"dockerImageVersionId":30732,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import time\nstart_time = time.time()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%python\n\n# Create test index 2500x50\n\nimport polars as pl\nimport whoosh_utils\nimport random\nimport os\n\ndata_dir = \"/kaggle/input/uspto-explainable-ai/\"\noutput_dir = \"/kaggle/working/\"\n\nnn_df = pl.read_csv(\"/kaggle/input/uspto-explainable-ai/test.csv\")\n# create directory to store the validation index in\nif not os.path.exists(os.path.join(output_dir, \"validation_index\")):\n    os.makedirs(os.path.join(output_dir, \"validation_index\"))\n\ntargets = nn_df[:,1:].to_numpy().ravel()\nfinal_targets = set(targets)\n\nprint(len(final_targets))\n\np_meta = pl.read_parquet(os.path.join(data_dir, \"patent_metadata.parquet\"), columns=[\"publication_number\", \"publication_date\", \"cpc_codes\"])\np_meta = p_meta.filter(pl.col(\"publication_number\").is_in(final_targets))\n\np_meta = p_meta.with_columns(pl.col(\"publication_date\").dt.year().alias(\"year\"))\np_meta = p_meta.with_columns(pl.col(\"publication_date\").dt.month().alias(\"month\"))\n\n\n# generate the documents that will go in the index\ndocuments = list()\nfor (year, month), meta_df in p_meta.group_by([\"year\", \"month\"]):\n    meta_df = meta_df.with_columns(pl.col(\"cpc_codes\").list.join(\" \"))\n    \n    try:\n        patents = pl.read_parquet(os.path.join(data_dir, f\"patent_data/{year}_{month}.parquet\"))\n        patents = patents.filter(pl.col(\"publication_number\").is_in(meta_df[\"publication_number\"]))\n        for i in range(meta_df.shape[0]):\n            d = dict()\n            p = patents.filter(pl.col(\"publication_number\") == meta_df[i, \"publication_number\"])\n            if p.shape[0] > 0:\n                d[\"publication_number\"] = p[0, \"publication_number\"]\n                d[\"title\"] = p[0, \"title\"]\n                d[\"abstract\"] = p[0, \"abstract\"]\n                d[\"claims\"] = p[0, \"claims\"]\n                d[\"description\"] = p[0, \"description\"]\n                d[\"cpc\"] = meta_df[i, \"cpc_codes\"]\n                documents.append(d)\n    except:\n        continue\n\n    \n    del patents\n    \nwhoosh_utils.create_index(os.path.join(output_dir, \"validation_index\"), documents)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:01:19.181296Z","iopub.execute_input":"2024-07-23T07:01:19.181676Z","iopub.status.idle":"2024-07-23T07:20:05.985753Z","shell.execute_reply.started":"2024-07-23T07:01:19.181644Z","shell.execute_reply":"2024-07-23T07:20:05.983174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport pickle\nimport re\nfrom collections import defaultdict\nfrom itertools import combinations\nfrom pathlib import Path\nfrom pprint import pprint\n\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport whoosh_utils\nimport whoosh\n\nfrom ortools.linear_solver import pywraplp\nfrom tqdm import tqdm\n\nINPUT_DIR = Path(\"../input/\")\nVALID = False\nDEBUG = False\nELEMENTS = [\n    \"cpc\",\n    \"comb_cpc\",\n    \"title_word\",\n    \"abstract_word\",\n    \"claims_word\",\n    \"description_word\",\n    \"title_bigram\",\n    \"abstract_bigram\",\n    \"claims_bigram\",\n    \"title_trigram\",\n    # \"abstract_trigram\",\n]\n\nif VALID:\n    nn_df = pl.read_csv(INPUT_DIR / \"uspto-explainable-ai-validation-index/neighbors_small.csv\")\nelse:\n    nn_df = pl.read_csv(INPUT_DIR / \"uspto-explainable-ai/test.csv\")\n\nif DEBUG:\n    nn_df = nn_df.sample(10, seed=42)\n\nmeta_df = pl.read_parquet(INPUT_DIR / \"uspto-explainable-ai/patent_metadata.parquet\")\nmeta_df = meta_df.filter(pl.col(\"publication_number\").is_in(nn_df.to_numpy().ravel()))\nmeta_df = meta_df.with_columns(year=meta_df[\"publication_date\"].dt.year(), month=meta_df[\"publication_date\"].dt.month())\n\n\nNUMBER_REGEX = re.compile(r\"^(\\d+|\\d{1,3}(,\\d{3})*)(\\.\\d+)?$\")\n# fmt: off\nBRS_STOPWORDS = [\n    \"an\", \"are\", \"by\", \"for\", \"if\", \"into\", \"is\", \"no\", \"not\", \"of\", \"on\", \"such\", \"that\", \"the\", \"their\", \"then\", \"there\", \"these\", \"they\", \"this\", \"to\", \"was\", \"will\",  # noqa\n]\n# fmt: on\n\n#valid_index = whoosh_utils.load_index(INPUT_DIR / \"uspto-explainable-ai-validation-index/validation/validation_index\")\n\n# testの2500x50 index\nvalid_index = whoosh_utils.load_index(\"./validation_index\")\n\nsearcher = whoosh_utils.get_searcher(valid_index)\nqp = whoosh_utils.get_query_parser()\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\ncustom_analyzer = whoosh.analysis.StandardAnalyzer(stoplist=BRS_STOPWORDS) | NumberFilter()\ncustom_analyzer_without_number_filter = whoosh.analysis.StandardAnalyzer(stoplist=BRS_STOPWORDS)\n\ndef extract_ngrams(text, n):\n    if n == 1:\n        tokens = [t.text for t in custom_analyzer(text)]\n        return list(set(tokens))\n\n    tokens = [t.text for t in custom_analyzer_without_number_filter(text)]\n    ngrams = list()\n    for i in range(len(tokens) - n + 1):\n        ngram = tokens[i : i + n]\n        # Numbers cannot be combined before and after for phrase search, so skip them.\n        if any(NUMBER_REGEX.match(token) for token in ngram):\n            continue\n        ngrams.append(tuple(ngram))\n\n    return list(set(ngrams))\n\n\ndetd_bad_words = set(pickle.load(open(INPUT_DIR / \"uspto-detd-bad-words/detd_bad_words.pickle\", \"rb\")))\n\n\ndef create_pnum_to_element(meta_df, element_name, global_counter=None):\n    pnum_to_element = defaultdict(set)\n    element_to_pnum = defaultdict(set)\n    if \"cpc\" in element_name:\n        col = element_name\n    else:\n        col, ngram = element_name.split(\"_\")\n        if ngram == \"word\":\n            ngram = 1\n        elif ngram == \"bigram\":\n            ngram = 2\n        elif ngram == \"trigram\":\n            ngram = 3\n        else:\n            raise ValueError(f\"Invalid ngram: {ngram}\")\n\n    if col == \"cpc\":\n        for row in tqdm(meta_df.iter_rows(named=True), total=meta_df.shape[0]):\n            elements = row[\"cpc_codes\"]\n            for element in elements:\n                pnum_to_element[row[\"publication_number\"]].add(element)\n                element_to_pnum[element].add(row[\"publication_number\"])\n        return pnum_to_element, element_to_pnum\n\n    if col == \"comb_cpc\":\n        for row in tqdm(meta_df.iter_rows(named=True), total=meta_df.shape[0]):\n            cpcs = row[\"cpc_codes\"]\n            elements = [\" \".join(c) for c in list(combinations(sorted(cpcs), 2))]\n            for element in elements:\n                pnum_to_element[row[\"publication_number\"]].add(element)\n                element_to_pnum[element].add(row[\"publication_number\"])\n        return pnum_to_element, element_to_pnum\n\n    assert ngram is not None\n    group_count = meta_df.select([\"year\", \"month\"]).n_unique()\n    for (year, month), grp in tqdm(meta_df.group_by([\"year\", \"month\"], maintain_order=True), total=group_count):\n        patents = pl.read_parquet(\n            INPUT_DIR / f\"uspto-explainable-ai/patent_data/{year}_{month}.parquet\",\n            columns=[\"publication_number\", col],\n        )\n        patents = patents.filter(pl.col(\"publication_number\").is_in(grp[\"publication_number\"]))\n        for row in patents.iter_rows(named=True):\n            elements = extract_ngrams(row[col], ngram)\n            for element in elements:\n                if col == \"description\" and element in detd_bad_words:\n                    continue \n                if global_counter is not None:\n                    if element not in global_counter:\n                        continue\n                pnum_to_element[row[\"publication_number\"]].add(element)\n                element_to_pnum[element].add(row[\"publication_number\"])\n    return pnum_to_element, element_to_pnum\n\n\ndef create_candidate_list(term_to_pnum, global_counter, kind, fp_allowance=5):\n    cand_list = []\n    for term in list(term_to_pnum.keys()):\n        covers = term_to_pnum[term]\n        if global_counter[term] == 0:\n            # In the case of 0, it is excluded due to high frequency\n            continue\n        fp = global_counter[term] - len(covers)\n        if fp > fp_allowance:\n            continue\n\n        if kind == \"cpc\":\n            cand = f'cpc:\"{term}\"'\n        if kind == \"comb_cpc\":\n            cand = \"\".join([f'cpc:\"{t}\"' for t in term.split(\" \")]) \n        if kind == \"title_word\":\n            cand = f'ti:\"{term}\"'\n        if kind == \"abstract_word\":\n            cand = f'ab:\"{term}\"'\n        if kind == \"claims_word\":\n            cand = f'clm:\"{term}\"'\n        if kind == \"description_word\":\n            cand = f'detd:\"{term}\"'\n        if kind == \"title_bigram\" or kind == \"title_trigram\":\n            cand = f'ti:\"{\"@\".join(term)}\"'\n        if kind == \"abstract_bigram\" or kind == \"abstract_trigram\":\n            cand = f'ab:\"{\"@\".join(term)}\"'\n        if kind == \"claims_bigram\":\n            cand = f'clm:\"{\"@\".join(term)}\"'\n\n        cand_list.append(\n            {\n                \"candidate\": cand,\n                \"tp_covers\": list(covers),\n                \"tp_count\": len(covers),\n                \"fp_count\": fp,\n                \"token\": len(cand.split(\" \")),\n                \"kind\": kind,\n            }\n        )\n    return cand_list\n\n\nglobal_counter_version = \"uspto-global-counters-limit30\"\n\n\ndef load_pnum_to_element(element_name, global_counter=None):\n    # if VALID:\n    #     path1 = Path(INPUT_DIR / f\"uspto-pnum-to-element/validation_pnum_to_{element_name}.pkl\")\n    #     path2 = Path(INPUT_DIR / f\"uspto-pnum-to-element/validation_{element_name}_to_pnum.pkl\")\n    #     print(f\"Loading validation_pnum_to_{element_name}...\")\n    #     pnum_to_element = pickle.load(open(path1, \"rb\"))\n    #     element_to_pnum = pickle.load(open(path2, \"rb\"))\n    # else:\n    path1 = Path(f\"pnum_to_{element_name}.pkl\")\n    path2 = Path(f\"{element_name}_to_element.pkl\")\n\n    if path1.exists() and path2.exists():\n        print(f\"Loading pnum_to_{element_name}...\")\n        pnum_to_element = pickle.load(open(path1, \"rb\"))\n        element_to_pnum = pickle.load(open(path2, \"rb\"))\n    else:\n        print(f\"Creating pnum_to_{element_name}...\")\n        if \"bigram\" not in element_name and \"trigram\" not in element_name:\n            global_counter = None\n        pnum_to_element, element_to_pnum = create_pnum_to_element(meta_df, element_name, global_counter=global_counter)\n        pickle.dump(pnum_to_element, open(path1, \"wb\"))\n        pickle.dump(element_to_pnum, open(path2, \"wb\"))\n    return pnum_to_element, element_to_pnum\n\n\ndef create_candidate_dict(element_name):\n    global_counter = pickle.load(open(INPUT_DIR / f\"{global_counter_version}/global_{element_name}_counter.pkl\", \"rb\"))\n    pnum_to_element, element_to_pnum = load_pnum_to_element(element_name, global_counter=global_counter)\n\n    candidate_dict = {}\n    for nn_row in tqdm(nn_df.iter_rows(), total=len(nn_df)):\n        pubunum = nn_row[0]\n        pnums = nn_row[1:]\n        nn_element_to_pnum = defaultdict(set)\n        for pnum in pnums:\n            for element in pnum_to_element[pnum]:\n                nn_element_to_pnum[element].add(pnum)\n\n        candidate = create_candidate_list(nn_element_to_pnum, global_counter, element_name)\n        candidate_dict[pubunum] = candidate\n\n    del global_counter, pnum_to_element, element_to_pnum\n    gc.collect()\n    return candidate_dict\n\n\ndef greedy_covering(candidates, target_patents):\n    selected_candidates = []\n    covered_patents = set()\n    token_cnt = 0\n\n    while len(covered_patents) < len(target_patents):\n        best_candidate = None\n        best_cover = set()\n        best_cost = float(\"inf\")\n        best_value = 0\n\n        for row in candidates:\n            if row[\"candidate\"] in selected_candidates:\n                continue\n\n            candidate_patents = set(row[\"tp_covers\"])\n            new_cover = candidate_patents - covered_patents\n\n            new_value = (len(new_cover) - row[\"fp_count\"] / 10) / max(row[\"token\"], 1)\n            new_cost = max(row[\"token\"], 1)\n\n            if len(new_cover) > 0 and new_value > best_value and token_cnt + new_cost + 1 <= 51:\n                best_candidate = row[\"candidate\"]\n                best_cover = new_cover\n                best_cost = new_cost\n                best_value = new_value\n\n        if best_candidate is not None:\n            selected_candidates.append(best_candidate)\n            covered_patents.update(best_cover)\n            token_cnt += best_cost + 1\n        else:\n            break\n\n    return \" OR \".join(selected_candidates), covered_patents, len(covered_patents), token_cnt\n\n\ndef cand_lis_to_query(candidate_lis):\n    qlis = []\n    for c in candidate_lis:\n        qlis.append(\"(\" + c[\"candidate\"] + \")\")\n        \n    return \" OR \".join(qlis)\n\n\ndef solve_weighted_set_cover(subsets, costs, penalties):\n    all_elements = set(e for subset in subsets for e in subset)\n\n    solver = pywraplp.Solver.CreateSolver(\"SCIP\")\n    if not solver:\n        return None\n\n    x = []\n    for i in range(len(subsets)):\n        x.append(solver.BoolVar(f\"x[{i}]\"))\n\n    z = {}\n    for e in all_elements:\n        z[e] = solver.BoolVar(f\"z[{e}]\")\n\n    for e in all_elements:\n        solver.Add(z[e] <= sum(x[i] for i in range(len(subsets)) if e in subsets[i]))\n\n    solver.Add(sum(costs[i] * x[i] for i in range(len(subsets))) + sum(x[i] for i in range(len(subsets))) <= 51)\n\n    objective = solver.Objective()\n    for e in all_elements:\n        objective.SetCoefficient(z[e], 10)\n    for i in range(len(subsets)):\n        objective.SetCoefficient(x[i], -penalties[i])\n    objective.SetMaximization()\n\n    status = solver.Solve()\n\n    if status == pywraplp.Solver.OPTIMAL:\n        selected_subsets = [i for i in range(len(subsets)) if x[i].solution_value() == 1]\n        total_cost = sum(costs[i] for i in selected_subsets)\n        covered_elements = set(e for i in selected_subsets for e in subsets[i])\n        return selected_subsets, total_cost, covered_elements\n    else:\n        return None, None, None\n\n\ndef solver_covering(candidates):\n    d = defaultdict(int)\n\n    subsets = []\n    costs = []\n    penalties = []\n\n    now = 0\n    for cand in candidates:\n        tmpset = set()\n        for tp in cand[\"tp_covers\"]:\n            if tp not in d:\n                d[tp] = now\n                now += 1\n\n            tmpset.add(d[tp])\n        subsets.append(tmpset)\n        costs.append(cand[\"token\"])\n        penalties.append(cand[\"fp_count\"])\n\n    selected_subsets, total_cost, covered_elements = solve_weighted_set_cover(subsets, costs, penalties)\n    if selected_subsets is None:\n        return \"\", set()\n\n    cand_lis = [candidates[ss] for ss in selected_subsets]\n\n    scp_query = cand_lis_to_query(cand_lis)\n\n    return scp_query, covered_elements\n\n\n# Candidates from the global counter\nfor element in ELEMENTS:\n    candidate_dict = create_candidate_dict(element)\n    pickle.dump(candidate_dict, open(f\"{element}_candidate.pkl\", \"wb\"))\n\ncandidate_dict = {}\nfor element in ELEMENTS:\n    print(f\"Gathering {element} candidate...\")\n    _candidate_dict = pickle.load(open(f\"{element}_candidate.pkl\", \"rb\"))\n    for key in set(candidate_dict) | set(_candidate_dict): \n        candidate_dict[key] = candidate_dict.get(key, []) + _candidate_dict.get(key, [])\n\n\ndef load_all_pnum_to_element():\n    all_pnum_to_element = {}\n    all_element_to_pnum = {}\n\n    # cpc\n    pnum_to_element, element_to_pnum = load_pnum_to_element(\"cpc\", None)\n    all_pnum_to_element[\"cpc\"] = pnum_to_element\n    all_element_to_pnum[\"cpc\"] = element_to_pnum\n\n    # ti\n    pnum_to_element, element_to_pnum = load_pnum_to_element(\"title_word\", None)\n    all_pnum_to_element[\"ti\"] = pnum_to_element\n    all_element_to_pnum[\"ti\"] = element_to_pnum\n\n    # ab\n    pnum_to_element, element_to_pnum = load_pnum_to_element(\"abstract_word\", None)\n    all_pnum_to_element[\"ab\"] = pnum_to_element\n    all_element_to_pnum[\"ab\"] = element_to_pnum\n\n    # clm\n    pnum_to_element, element_to_pnum = load_pnum_to_element(\"claims_word\", None)\n    all_pnum_to_element[\"clm\"] = pnum_to_element\n    all_element_to_pnum[\"clm\"] = element_to_pnum\n\n    # detd\n    pnum_to_element, element_to_pnum = load_pnum_to_element(\"description_word\", None)\n    all_pnum_to_element[\"detd\"] = pnum_to_element\n    all_element_to_pnum[\"detd\"] = element_to_pnum\n\n    return all_pnum_to_element, all_element_to_pnum\n\n\nprint(\"Loading all pnum to element...\")\nall_pnum_to_element, all_element_to_pnum = load_all_pnum_to_element()\n\n\ndef get_cand_lis(pnum_lis, idf_min, elements):\n    cand_lis = []\n    for elem in elements:\n        pnum_to_element = all_pnum_to_element[elem]\n\n        word_list = pnum_to_element[pnum_lis[0]].copy()\n        for key in pnum_lis[1:]:\n            word_list &= pnum_to_element[key]\n\n        for w in word_list:\n            idf = searcher.idf(elem, w)\n            if idf < idf_min:\n                continue\n            cand_lis.append((idf, elem, w))\n\n    return cand_lis\n\n\ndef check_query_match(cand_lis, neighbor50):\n    tp_covers = set()\n\n    for pnum in neighbor50:\n        is_ok = True\n        for _, elem, word in cand_lis:\n            pnum_to_element = all_pnum_to_element[elem]\n            if word not in pnum_to_element[pnum]:\n                is_ok = False\n                break\n        if is_ok:\n            tp_covers.add(pnum)\n\n    first = True\n    for _, elem, word in cand_lis:\n        element_to_pnum = all_element_to_pnum[elem]\n        if first:\n            fp_locals = element_to_pnum[word] - neighbor50\n            first = False\n        else:\n            fp_locals = fp_locals & (element_to_pnum[word] - neighbor50)\n            if len(fp_locals) == 0:\n                break\n\n    return tp_covers, fp_locals\n\n\ndef get_cand(pnum_lis, neighbor50, elements):\n    word_idf_min = 4.0  # Do not use words with an IDF value below this.\n    sum_idf_min = 80.0  # The total IDF value of the created query must meet the minimum threshold\n    query_max = 20  # Connect up to this number of terms with AND\n\n    cand_lis = get_cand_lis(pnum_lis, word_idf_min, elements)  # list of (idf, elem, word)\n\n    # Select up to query_max terms in descending order of IDF.\n    cand_lis.sort(reverse=True)\n    cand_lis = cand_lis[:query_max]\n\n    idf_sum = sum([tup[0] for tup in cand_lis])\n    if idf_sum < sum_idf_min:\n        return None\n\n    tp_covers, fp_locals = check_query_match(cand_lis, neighbor50)\n\n    q_lis = []\n    for _, elm, word in cand_lis:\n        q_lis.append(f'{elm}:\"{word}\"')\n    query = \"\".join(q_lis)\n\n    return tp_covers, fp_locals, query, idf_sum\n\n\ndef create_fp0_cand(neighbor50, kind_lis, c):\n    cand_list = []\n\n    # 50 c 2\n    for pnum_lis in combinations(neighbor50, c):\n        ret = get_cand(pnum_lis, neighbor50, kind_lis)\n        if ret is None:\n            continue\n\n        tp_covers, fp_locals, query, idf_sum = ret\n        if idf_sum < 110:\n            fp = 1 + len(fp_locals)\n        else:\n            fp = len(fp_locals)\n\n        cand_list.append(\n            {\n                \"candidate\": \"(\" + query + \")\",\n                \"tp_covers\": tp_covers,\n                \"fp_locals\": fp_locals,\n                \"tp_count\": len(tp_covers),\n                \"fp_count\": fp,\n                \"token\": 1,\n                \"kind\": \"fp0\",\n            }\n        )\n    return cand_list\n\n\ndef create_compressed_candidates(candidates):\n    tp_covers_dict = defaultdict(list)\n\n    for cand in candidates:\n        tmp = (cand[\"fp_count\"],cand[\"candidate\"])\n        tp_covers_dict[tuple(cand[\"tp_covers\"])].append(tmp)\n\n    ret = []\n    \n    for tp_covers,v in tp_covers_dict.items():\n        v.sort()\n        \n        fp_min = 10**9\n        candidate_lis = []\n        \n        cnt=0\n        for e in v:\n            fp_min = min(fp_min, e[0])\n            candidate_lis.append(e[1])\n            cnt+=1\n            if cnt>10:\n                break\n\n        tmp = {\n                \"candidate\":  \"\".join(candidate_lis),\n                \"tp_covers\": tp_covers,\n                \"tp_count\": len(tp_covers),\n                \"fp_count\": fp_min,\n                \"token\": 1,\n                \"kind\": \"new_global_counter\",\n            }\n        ret.append(tmp)\n    return ret\n","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:20:32.521956Z","iopub.execute_input":"2024-07-23T07:20:32.522365Z","iopub.status.idle":"2024-07-23T07:48:58.288717Z","shell.execute_reply.started":"2024-07-23T07:20:32.522336Z","shell.execute_reply":"2024-07-23T07:48:58.285193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"backup_query = \"ti:device\"\nvalid_pnum = []\ngts = []\nqueries = []\ncand_len_lis = []\ncnt = 0\nprint(\"Creating queries from candidates...\")\nfor nn_row in tqdm(nn_df.iter_rows(), total=len(nn_df)):\n    candidates = candidate_dict[nn_row[0]]\n    candidates = create_compressed_candidates(candidates)\n\n    neighbor50 = set(nn_row[1:])\n    cand_list_cpc = create_fp0_cand(neighbor50, [\"ti\", \"ab\", \"clm\", \"cpc\"], 2)\n    cand_list_detd = create_fp0_cand(neighbor50, [\"ti\", \"ab\", \"clm\", \"cpc\", \"detd\"], 2)\n\n    query_cpc, covered_patents = solver_covering(candidates + cand_list_cpc)\n    query_detd, covered_patents = solver_covering(candidates + cand_list_detd)\n\n    if len(query_detd) == 0 or len(query_detd) > 10_000:\n        # If the DETD query exceeds 10,000, the CPC query must be used.\n        query = query_cpc\n    elif len(cand_list_cpc) > 90:\n        # If many query candidates are created, there is no need to use DETD candidates.\n        query = query_cpc\n    else:\n        # Since not many query candidates are created, use DETD candidates as well.\n        query = query_detd\n\n    queries.append(query)\n    cand_len_lis.append(len(cand_list_cpc))\n    # org_cand_len.append(len(candidates))\n\n    valid_pnum.append(nn_row[0])\n    gts.append(nn_row[1:])\n","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:47.83582Z","iopub.execute_input":"2024-07-23T07:57:47.837011Z","iopub.status.idle":"2024-07-23T07:57:50.01263Z","shell.execute_reply.started":"2024-07-23T07:57:47.836964Z","shell.execute_reply":"2024-07-23T07:57:50.011292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def ap50(preds, labels):\n    ox_lis = []\n    precisions = list()\n    n_found = 0\n    for e, i in enumerate(range(50)):\n        if i < len(preds):\n            if preds[i] in labels:\n                n_found += 1\n                ox_lis.append(\"○\")\n            else:\n                ox_lis.append(\"✘\")\n        precisions.append(n_found / (e + 1))\n    return sum(precisions) / 50, ox_lis","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:50.014749Z","iopub.execute_input":"2024-07-23T07:57:50.015261Z","iopub.status.idle":"2024-07-23T07:57:50.023124Z","shell.execute_reply.started":"2024-07-23T07:57:50.015219Z","shell.execute_reply":"2024-07-23T07:57:50.021692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ap_lis = []\n\nfor i in range(len(valid_pnum)):\n    pnum = valid_pnum[i]\n    query = queries[i]\n    \n    ret = whoosh_utils.execute_query(query, qp, searcher)\n    \n    ap, ox = ap50(ret, gts[i])\n    ap_lis.append(ap)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:50.084369Z","iopub.execute_input":"2024-07-23T07:57:50.085241Z","iopub.status.idle":"2024-07-23T07:57:51.043648Z","shell.execute_reply.started":"2024-07-23T07:57:50.085202Z","shell.execute_reply":"2024-07-23T07:57:51.042568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\norg_result_df = pd.DataFrame()\norg_result_df[\"publication_number\"] = valid_pnum\norg_result_df[\"ap\"] = ap_lis\norg_result_df[\"query\"] = queries\n\n# ap昇順でソート\nresult_df = org_result_df.sort_values(\"ap\").copy().reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:52.65089Z","iopub.execute_input":"2024-07-23T07:57:52.651289Z","iopub.status.idle":"2024-07-23T07:57:52.661725Z","shell.execute_reply.started":"2024-07-23T07:57:52.651259Z","shell.execute_reply":"2024-07-23T07:57:52.660424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"result_df","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:53.181006Z","iopub.execute_input":"2024-07-23T07:57:53.182094Z","iopub.status.idle":"2024-07-23T07:57:53.194294Z","shell.execute_reply.started":"2024-07-23T07:57:53.182052Z","shell.execute_reply":"2024-07-23T07:57:53.193163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_pnum = result_df[\"publication_number\"].values\nap_lis = result_df[\"ap\"].values\nqueries = result_df[\"query\"].values","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:53.622474Z","iopub.execute_input":"2024-07-23T07:57:53.622911Z","iopub.status.idle":"2024-07-23T07:57:53.628747Z","shell.execute_reply.started":"2024-07-23T07:57:53.622872Z","shell.execute_reply":"2024-07-23T07:57:53.627581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_row(pnum):\n    nn_row = nn_df.filter(pl.col(\"publication_number\") == pnum).row(0)\n    return nn_row","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:54.038745Z","iopub.execute_input":"2024-07-23T07:57:54.039186Z","iopub.status.idle":"2024-07-23T07:57:54.044637Z","shell.execute_reply.started":"2024-07-23T07:57:54.039151Z","shell.execute_reply":"2024-07-23T07:57:54.043366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fp0_cand(cand_lis):\n    ret = []\n    for cand in cand_lis:\n        if cand[\"fp_count\"]==0:\n            ret.append(cand)\n    return ret","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:54.246123Z","iopub.execute_input":"2024-07-23T07:57:54.246534Z","iopub.status.idle":"2024-07-23T07:57:54.252116Z","shell.execute_reply.started":"2024-07-23T07:57:54.246499Z","shell.execute_reply":"2024-07-23T07:57:54.25103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"backup_query = \"ti:device\"\nqueries_fp0 = []\ncnt = 0\nprint(\"Creating queries from candidates...\")\n\n\nfor i in range(len(valid_pnum)):\n    pnum = valid_pnum[i]\n    \n    nn_row = get_row(pnum)\n    \n    neighbor50 = set(nn_row[1:])\n    \n    candidates = candidate_dict[nn_row[0]]\n    candidates = create_compressed_candidates(candidates)\n\n    \n    cand_list_cpc = create_fp0_cand(neighbor50, [\"ti\", \"ab\", \"clm\", \"cpc\"], 2)\n    cand_list_detd = create_fp0_cand(neighbor50, [\"ti\", \"ab\", \"clm\", \"cpc\", \"detd\"], 2)\n\n    concat_cpc = candidates + cand_list_cpc\n    concat_detd = candidates + cand_list_detd\n    \n    concat_cpc = fp0_cand(concat_cpc)\n    concat_detd = fp0_cand(concat_detd)\n    \n    query_cpc,  covered_patents = solver_covering(concat_cpc)\n    query_detd, covered_patents = solver_covering(concat_detd)\n    \n\n    if len(query_detd) == 0 or len(query_detd) > 10_000:\n        # If the DETD query exceeds 10,000, the CPC query must be used.\n        query = query_cpc\n    elif len(cand_list_cpc) > 90:\n        # If many query candidates are created, there is no need to use DETD candidates.\n        query = query_cpc\n    else:\n        # Since not many query candidates are created, use DETD candidates as well.\n        query = query_detd\n\n\n    ret = whoosh_utils.execute_query(query, qp, searcher)\n    ap, ox = ap50(ret, nn_row[1:])\n    \n    if ap > ap_lis[i]:\n        queries[i] = query\n        \n    elapsed_time = time.time() - start_time\n    if elapsed_time > 8.5 * 60 * 60:\n        break","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:54.440582Z","iopub.execute_input":"2024-07-23T07:57:54.441003Z","iopub.status.idle":"2024-07-23T07:57:56.897846Z","shell.execute_reply.started":"2024-07-23T07:57:54.440971Z","shell.execute_reply":"2024-07-23T07:57:56.89667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"time.time() - start_time","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:56.899862Z","iopub.execute_input":"2024-07-23T07:57:56.900242Z","iopub.status.idle":"2024-07-23T07:57:56.907986Z","shell.execute_reply.started":"2024-07-23T07:57:56.900203Z","shell.execute_reply":"2024-07-23T07:57:56.90681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"result_df = pd.DataFrame()\nresult_df[\"publication_number\"] = valid_pnum\nresult_df[\"query\"] = queries","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:56.909538Z","iopub.execute_input":"2024-07-23T07:57:56.910018Z","iopub.status.idle":"2024-07-23T07:57:56.921032Z","shell.execute_reply.started":"2024-07-23T07:57:56.909954Z","shell.execute_reply":"2024-07-23T07:57:56.919787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"result_df","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:56.933185Z","iopub.execute_input":"2024-07-23T07:57:56.933606Z","iopub.status.idle":"2024-07-23T07:57:56.945356Z","shell.execute_reply.started":"2024-07-23T07:57:56.933573Z","shell.execute_reply":"2024-07-23T07:57:56.944061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"org_result_df","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:57:57.699158Z","iopub.execute_input":"2024-07-23T07:57:57.7002Z","iopub.status.idle":"2024-07-23T07:57:57.711879Z","shell.execute_reply.started":"2024-07-23T07:57:57.700161Z","shell.execute_reply":"2024-07-23T07:57:57.710777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_df = pd.merge(org_result_df[[\"publication_number\"]], result_df, on=\"publication_number\",how=\"left\")","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:58:00.494164Z","iopub.execute_input":"2024-07-23T07:58:00.495068Z","iopub.status.idle":"2024-07-23T07:58:00.504547Z","shell.execute_reply.started":"2024-07-23T07:58:00.495025Z","shell.execute_reply":"2024-07-23T07:58:00.503317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_df","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:58:00.668679Z","iopub.execute_input":"2024-07-23T07:58:00.669574Z","iopub.status.idle":"2024-07-23T07:58:00.680236Z","shell.execute_reply.started":"2024-07-23T07:58:00.669532Z","shell.execute_reply":"2024-07-23T07:58:00.678907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nif VALID:\n    print(\"Validation, executing queries...\")\n    preds = []\n    for i in tqdm(range(len(queries))):\n        query = queries[i]\n        try:\n            ret = whoosh_utils.execute_query(query, qp, searcher)\n        except Exception:\n            print(\">>>\", query)\n            ret = whoosh_utils.execute_query(backup_query, qp, searcher)\n        preds.append(ret)\n\n    ap_lis = []\n    ox_lis = []\n    for i in range(len(preds)):\n        pred = preds[i]\n        ap, ox = ap50(pred, gts[i])\n        ap_lis.append(ap)\n        ox_lis.append(\"\".join(ox))\n    print(\"Mean AP\", np.mean(ap_lis))\n    df = pd.DataFrame()\n    df[\"pnum\"] = valid_pnum\n    df[\"ap\"] = ap_lis\n    df[\"ox\"] = ox_lis\n    df[\"cand_len_cpc\"] = cand_len_lis\n    # df[\"cand_len_detd\"] = cand_len_detd\n    # df[\"org_cand_len\"] = org_cand_len\n    df.to_csv(\"valid_result.csv\", index=None)\nelse:\n    #nn_df = nn_df.with_columns(pl.Series(\"query\", queries))\n    print(final_df[[\"publication_number\", \"query\"]].head())\n    final_df[[\"publication_number\", \"query\"]].to_csv(\"submission.csv\",index=None)\n","metadata":{"execution":{"iopub.status.busy":"2024-07-23T07:58:23.724961Z","iopub.execute_input":"2024-07-23T07:58:23.725362Z","iopub.status.idle":"2024-07-23T07:58:23.744024Z","shell.execute_reply.started":"2024-07-23T07:58:23.725331Z","shell.execute_reply":"2024-07-23T07:58:23.742679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}