{"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":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 gc\nimport itertools\nimport pickle\nimport re\nfrom collections import defaultdict\nfrom pathlib import Path\n\nimport numpy as np\nimport polars as pl\nimport whoosh_utils\nimport whoosh\n\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)\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\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\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        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\ndef create_pnum_to_element(meta_df, global_counter, element_name):\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    pnum_to_element = defaultdict(set)\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                if element in global_counter:\n                    pnum_to_element[row[\"publication_number\"]].add(element)\n        return pnum_to_element\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(itertools.combinations(sorted(cpcs), 2))]\n            for element in elements:\n                if element in global_counter:\n                    pnum_to_element[row[\"publication_number\"]].add(element)\n        return pnum_to_element\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 element in global_counter:\n                    pnum_to_element[row[\"publication_number\"]].add(element)\n\n    return pnum_to_element\n\n\ndef create_candidate_list(term_to_pnum, global_counter, kind, fp_allowance=10):\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            # 0 の場合は頻度が多いので除外されている term\n            continue\n        fp = global_counter[term] - len(covers)\n        if fp > fp_allowance:  # どれだけ FP を許容するかは調整可能\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 create_candidate_dict(element_name):\n    global_counter = pickle.load(open(INPUT_DIR / f\"{global_counter_version}/global_{element_name}_counter.pkl\", \"rb\"))\n    path = Path(INPUT_DIR / f\"uspto-pnum-to-element/validation_pnum_to_{element_name}.pkl\")\n    # path = Path(\"dummy\")\n    if VALID and path.exists():\n        print(f\"Loading pnum_to_{element_name}...\")\n        pnum_to_element = pickle.load(open(path, \"rb\"))\n    else:\n        print(f\"Creating pnum_to_{element_name}...\")\n        pnum_to_element = create_pnum_to_element(meta_df, global_counter, element_name)\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\n    gc.collect()\n    return candidate_dict\n\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\n        # 最も多くの未カバー特許をカバーする候補を選ぶ\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  # 1 は OR の分\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\nfrom ortools.linear_solver import pywraplp\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))) <= 50)\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 not tp 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\nfor element in ELEMENTS:\n    print(f\"Creating {element} candidate...\")\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\nif VALID:\n    valid_index = whoosh_utils.load_index(INPUT_DIR / \"uspto-explainable-ai-validation-index/validation/validation_index\")\n    searcher = whoosh_utils.get_searcher(valid_index)\n    qp = whoosh_utils.get_query_parser()\n\nbackup_query = \"ti:device\"\nvalid_pnum = []\npreds = []\ngts = []\nqueries = []\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    query, covered_patents = solver_covering(candidates)\n    \n    if len(query) == 0:\n        query = backup_query\n    queries.append(query)\n    if VALID:\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        while len(ret) < 50:\n            ret.append(\"\")\n\n        valid_pnum.append(nn_row[0])\n        preds.append(ret)\n        gts.append(nn_row[1:])\n\n\ndef ap50(preds, labels):\n    precisions = list()\n    n_found = 0\n    for e, i in enumerate(range(50)):\n        if i < len(preds) and preds[i] in labels:\n            n_found += 1\n        precisions.append(n_found / (e + 1))\n    return sum(precisions) / 50\n\n\nif VALID:\n    pred_num = []\n    ap_lis = []\n    for i in range(len(preds)):\n        pred = preds[i]\n        pred_num.append(len(pred))\n        pnum = valid_pnum[i]\n\n        ap = ap50(pred, gts[i])\n        ap_lis.append(ap)\n\n    print(\"Mean AP\", np.mean(ap_lis))\nelse:\n    nn_df = nn_df.with_columns(pl.Series(\"query\", queries))\n    print(nn_df[[\"publication_number\", \"query\"]].head())\n    nn_df[[\"publication_number\", \"query\"]].write_csv(\"submission.csv\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-02T14:23:28.020862Z","iopub.execute_input":"2024-07-02T14:23:28.021299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}