{"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":8479599,"sourceType":"datasetVersion","datasetId":4517815},{"sourceId":8984766,"sourceType":"datasetVersion","datasetId":5297580},{"sourceId":9026724,"sourceType":"datasetVersion","datasetId":4930666},{"sourceId":189607865,"sourceType":"kernelVersion"}],"dockerImageVersionId":30698,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import polars as pl\nimport pyarrow.dataset as ds\nimport pyarrow as pa\nimport os\nimport json\nfrom tqdm import tqdm\nimport numpy as np\nimport gc\nimport pickle\nimport csv\nimport time\nfrom collections import Counter, defaultdict\nfrom typing import List\nimport math\nimport itertools\nimport whoosh_utils\nimport heapq\nfrom scipy import stats\nimport re\nimport string\nfrom whoosh.searching import TimeLimit","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-23T15:55:10.986654Z","iopub.execute_input":"2024-07-23T15:55:10.987061Z","iopub.status.idle":"2024-07-23T15:55:50.922005Z","shell.execute_reply.started":"2024-07-23T15:55:10.987027Z","shell.execute_reply":"2024-07-23T15:55:50.920747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"root_dir = '/kaggle/input/uspto-explainable-ai'","metadata":{"execution":{"iopub.status.busy":"2024-07-23T15:55:50.924413Z","iopub.execute_input":"2024-07-23T15:55:50.925133Z","iopub.status.idle":"2024-07-23T15:55:50.934224Z","shell.execute_reply.started":"2024-07-23T15:55:50.925093Z","shell.execute_reply":"2024-07-23T15:55:50.931076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_fields = ['title', 'abstract', 'claims', 'description']\nall_fields_phrase_lf = []\n\nfor i in all_fields:\n    if i == 'description':\n        p_lf_1 = pl.scan_parquet(f'/kaggle/input/uspto-preprocessed/sh_filtered_phrases_{i}_1')\n        p_lf_2 = pl.scan_parquet(f'/kaggle/input/uspto-preprocessed/sh_filtered_phrases_{i}_2')\n        p_lf_3 = pl.scan_parquet(f'/kaggle/input/uspto-preprocessed/sh_filtered_phrases_{i}_3')\n        p_lf = pl.concat((p_lf_1, p_lf_2, p_lf_3))\n    else:\n        p_lf = pl.scan_parquet(f'/kaggle/input/uspto-preprocessed/sh_filtered_phrases_{i}')\n    all_fields_phrase_lf.append(p_lf)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T15:55:50.935922Z","iopub.execute_input":"2024-07-23T15:55:50.936736Z","iopub.status.idle":"2024-07-23T15:55:51.103821Z","shell.execute_reply.started":"2024-07-23T15:55:50.936690Z","shell.execute_reply":"2024-07-23T15:55:51.102479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_csv = '/kaggle/input/validation-25k-1/test_neighbors.csv'\ntest_csv = os.path.join(root_dir, 'test.csv')","metadata":{"execution":{"iopub.status.busy":"2024-07-23T15:55:51.179079Z","iopub.execute_input":"2024-07-23T15:55:51.179469Z","iopub.status.idle":"2024-07-23T15:55:51.184951Z","shell.execute_reply.started":"2024-07-23T15:55:51.179435Z","shell.execute_reply":"2024-07-23T15:55:51.183703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_ids = set()\nwith open(test_csv) as csvfile:  \n    data = csv.reader(csvfile)\n    row_count = 0\n    \n    for row in data:\n        if row_count == 0:\n            row_count += 1\n            continue\n        \n        for i in row[1:]:\n            all_ids.add(i)\n\nprint(len(all_ids))","metadata":{"execution":{"iopub.status.busy":"2024-07-23T15:55:51.186915Z","iopub.execute_input":"2024-07-23T15:55:51.187303Z","iopub.status.idle":"2024-07-23T15:55:51.203131Z","shell.execute_reply.started":"2024-07-23T15:55:51.187267Z","shell.execute_reply":"2024-07-23T15:55:51.201715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_fields_phrase_data = {}\n\nfor idx, i in enumerate(all_fields_phrase_lf):\n    df = i.filter(pl.col('publication_number').is_in(all_ids)).collect()\n    print(df)\n    \n    label = all_fields[idx]\n    print(label)\n    with open(f'/kaggle/input/uspto-preprocessed/{label}_vocab_phrases_keys.json') as f:\n        vocab = json.load(f)\n    \n    all_fields_phrase_data[label] = (dict(df.iter_rows()), vocab)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T15:55:51.205330Z","iopub.execute_input":"2024-07-23T15:55:51.205787Z","iopub.status.idle":"2024-07-23T16:02:09.019658Z","shell.execute_reply.started":"2024-07-23T15:55:51.205753Z","shell.execute_reply":"2024-07-23T16:02:09.017293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PRINT_TO_LOG = False\n\nif PRINT_TO_LOG:\n    debug_log = 'debug_log.txt'\n    with open(debug_log, 'w') as f:\n        pass\n\ndef lprint(debug_obj):\n    if PRINT_TO_LOG:\n        with open(debug_log, 'a') as f:\n            print(debug_obj, file=f)\n    else:\n        print(debug_obj)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T16:02:09.022183Z","iopub.execute_input":"2024-07-23T16:02:09.022862Z","iopub.status.idle":"2024-07-23T16:02:09.032496Z","shell.execute_reply.started":"2024-07-23T16:02:09.022806Z","shell.execute_reply":"2024-07-23T16:02:09.031140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ref. https://github.com/benhamner/Metrics/blob/master/Python/ml_metrics/average_precision.py\ndef apk(actual, predicted, k=50):\n    if not actual:\n        return 0.0\n\n    if len(predicted)>k:\n        predicted = predicted[:k]\n\n    score = 0.0\n    num_hits = 0.0\n\n    for i,p in enumerate(predicted):\n        # first condition checks whether it is valid prediction\n        # second condition checks if prediction is not repeated\n        if p in actual and p not in predicted[:i]:\n            num_hits += 1.0\n            score += num_hits / (i+1.0)\n    \n    return score / min(len(actual), k)\n\ndef mapk(actual, predicted, k=50):\n    return np.mean([apk(a,p,k) for a,p in zip(actual, predicted)])","metadata":{"execution":{"iopub.status.busy":"2024-07-23T16:02:09.034833Z","iopub.execute_input":"2024-07-23T16:02:09.035337Z","iopub.status.idle":"2024-07-23T16:02:09.057127Z","shell.execute_reply.started":"2024-07-23T16:02:09.035290Z","shell.execute_reply":"2024-07-23T16:02:09.055786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def ideal_apk_for_cover(cover):\n    actual = [i for i in range(50)]\n    predicted = [i for i in range(cover)] + [-1]*(50-cover)\n    \n    return apk(actual, predicted)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T16:02:09.063404Z","iopub.execute_input":"2024-07-23T16:02:09.064584Z","iopub.status.idle":"2024-07-23T16:02:09.071981Z","shell.execute_reply.started":"2024-07-23T16:02:09.064532Z","shell.execute_reply":"2024-07-23T16:02:09.070626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = 0\n\nquery_validator = whoosh_utils.QueryValidator()\nqp = whoosh_utils.get_query_parser()\ntrain_idx = whoosh_utils.load_index('/kaggle/input/validation-25k-1')\nsearcher = whoosh_utils.get_searcher(train_idx)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T16:02:09.073540Z","iopub.execute_input":"2024-07-23T16:02:09.074056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"queries = []\ncurr_cover_dist = []\nmax_cover_dist = []\napk_accuracy = []\napks = []\ntoken_efficiency = []\ntoken_usage = []\nquery_time = []\nmax_time = 0\ntotal_time = 0\ndefault_query = 'ti:device'\nalpha = 0.32\nweight_threshold = 0.012\ncomb_threshold = 0.27\n\nwith open(test_csv) as csvfile:  \n    data = csv.reader(csvfile)\n    row_count = 0\n    \n    for row in tqdm(data):\n        if row_count == 0:\n            row_count += 1\n            continue\n        \n        ids = row[1:]\n        targets = len(ids)\n        ranked_ids = np.array(ids)\n        all_combs = {}\n        query_string = default_query\n        \n        def process_sh_field_ranked(labels, threshold):\n            field_scores = defaultdict(float)\n            truth_table = defaultdict(int)\n            max_weights = defaultdict(float)\n            final_scores = {}\n            assignments = []\n            fields = []\n                \n            for label in labels:\n                source = all_fields_phrase_data[label]\n                fielddata = [source[0].get(i, []) for i in ids]\n                for idx, i in enumerate(fielddata):\n                    for j in range(0, len(i), 3):\n                        word = source[1][i[j]]\n                        rank = i[j+1]\n                        dfreq = i[j+2]\n                        field_scores[(label, word)] += alpha*(dfreq-rank+1)/(dfreq*(dfreq+1)/2.0)+(1.0-alpha)/dfreq\n                        truth_table[(label, word)] |= (1 << idx)\n                \n                for k, v in truth_table.items():\n                    if field_scores[k] > threshold and field_scores[k] > max_weights[v]:\n                        final_scores[v] = (v, k[0], k[1], field_scores[k])\n                        max_weights[v] = field_scores[k]\n                        \n            vals = list(final_scores.values())\n            \n            del fielddata, field_scores, truth_table, max_weights, final_scores\n            \n            return vals\n        \n        all_assignments = process_sh_field_ranked(['title', 'abstract', 'claims', 'description'], weight_threshold)\n        \n        def q_for_key(key):\n            key_val = key[1].replace(' ', '-')\n            if key[0] == 'cpc':\n                return 'cpc:'+key_val\n            elif key[0] == 'title':\n                return 'ti:'+key_val\n            elif key[0] == 'abstract':\n                return 'ab:'+key_val\n            elif key[0] == 'claims':\n                return 'clm:'+key_val\n            elif key[0] == 'description':\n                return 'detd:'+key_val\n            else:\n                print(f'key error {key}')\n        \n        onecover = 0\n        if all_assignments is not None:\n            for i in all_assignments:\n                onecover |= i[0]\n                all_combs[q_for_key((i[1], i[2]))] = (i[0], i[3])\n        \n        def count_ones(n):\n            count = 0\n            while n:\n                n &= n - 1\n                count += 1\n            return count\n        \n        if DEBUG >= 1:\n            max_onecover = count_ones(onecover)\n            max_cover_dist.append(max_onecover)\n            missing = [ranked_ids[i] for i in range(targets) if not (onecover & (1 << i))]\n            lprint(f'maximum cover is {max_onecover}')\n            lprint(missing)\n        \n        tokens_count = 0\n        all_cover = []\n        query_list = []\n        max_weight = 0.0\n        max_tokens = 49 # each term takes 2 tokens including operator\n        \n        while tokens_count < max_tokens and len(all_combs) > 0: \n            choice = None\n            chosen_cover = None\n            if DEBUG >= 1:\n                lprint(len(all_combs))\n            for k, v in all_combs.items():\n                assignment = v[0]\n                weight = v[1]\n                cover = all_cover.copy()\n                cover.append(assignment)\n                cover = np.array(cover).reshape(-1, 1)\n                cover = (cover & (1 << np.arange(targets))) > 0\n                query_weights = [i[2] for i in query_list]\n                query_weights.append(weight)\n                query_weights = np.array(query_weights)\n                weighted_sum = (np.minimum((cover * query_weights[:, np.newaxis]).sum(axis=0), comb_threshold)).sum()\n                if weighted_sum > max_weight:\n                    choice = (k, v[0], v[1])\n                    max_weight = weighted_sum\n                    chosen_cover = assignment\n                    if DEBUG >= 1:\n                        lprint(f'chosen {choice[0]} {choice[2]} max weight {max_weight}')\n            if choice is None:\n                if DEBUG >= 1:\n                    lprint(f'nothing chosen {max_weight}')\n                break\n            if DEBUG >= 1:\n                lprint(f'chosen {choice[0]} {choice[2]} max weight {max_weight}')\n            del all_combs[choice[0]]\n            all_cover.append(chosen_cover)\n            query_list.append(choice)\n            if len(query_list) > 0:\n                query_string = ' OR '.join([i[0] for i in query_list])\n                tokens_count = whoosh_utils.count_query_tokens(query_string)\n        \n        curr_cover = 0\n        for i in all_cover:\n            curr_cover |= i\n        \n        final_cover = count_ones(curr_cover)\n        lprint(f'final cover {final_cover}')\n        lprint(f'final token count {tokens_count}')\n        \n        try:\n            query_validator.validate_query(query_string)\n        except Exception as e:\n            lprint(f'Query validation failed {e}')\n            query_string = default_query\n        \n        lprint('attempting query')\n        lprint(query_string)\n#         lprint(qp.parse(query_string))\n\n        start = time.perf_counter()\n        try:\n            results = whoosh_utils.execute_query(query_string, qp, searcher, 8.0)\n            clock = time.perf_counter()-start\n        except TimeLimit as tl:\n            clock = time.perf_counter()-start\n            lprint(tl.args[0])\n            results = tl.args[1]\n            partial_results = [(i[0], i[1].decode('utf-8')) for i in results[1]]\n            remove_idxs = []\n            for idx, i in enumerate(query_list):\n                parts = i[0].split(':')\n                field = parts[0]\n                if field != 'detd':\n                    continue\n                term = parts[1]\n                terms = term.replace('\"', '').split('-')\n                for j in terms:\n                    if (field, j) not in partial_results:\n                        remove_idxs.append(idx)\n            query_list = [i for idx, i in enumerate(query_list) if idx not in remove_idxs]\n            query_string = ' OR '.join([i[0] for i in query_list])\n            lprint(query_string)\n            tokens_count = whoosh_utils.count_query_tokens(query_string)\n        except Exception as e:\n            clock = time.perf_counter()-start\n            lprint('resetting due to other error')\n            query_string = default_query\n            results = (None, None)\n        total_time += clock\n        \n        if DEBUG >= 1:\n            if clock > max_time:\n                max_time = clock\n            lprint(f'query took {clock}')\n            query_time.append(clock)\n            lprint(ids)\n            lprint(results)\n            if results[0] is not None:\n                results_docs = results[0]\n                score = apk(ids, results_docs)\n                apks.append(score)\n                lprint(f'{row_count} {row[0]} score is {score}')\n                curr_cover_dist.append(final_cover)\n                ideal_apk = ideal_apk_for_cover(final_cover)\n                lprint(f'ideal apk {ideal_apk} for cover {final_cover}')\n                if ideal_apk <= 0:\n                    apk_accuracy.append(0)\n                else:\n                    apk_accuracy.append(score/ideal_apk)\n                token_usage.append(tokens_count)\n                if tokens_count > 0:\n                    efficiency = count_ones(curr_cover)/tokens_count\n                    token_efficiency.append(efficiency)\n                else:\n                    token_efficiency.append(0)\n                lprint(results_docs)\n                del results_docs\n            else:\n                apks.append(0)\n                curr_cover_dist.append(0)\n                apk_accuracy.append(0)\n                token_usage.append(0)\n                token_efficiency.append(0)\n        \n        lprint(query_string)\n        queries.append(query_string)\n        \n        del ids, ranked_ids\n        del all_assignments, all_combs\n        del all_cover, query_list\n        del results\n        \n        row_count += 1\n#         if row_count > 10:\n#             break","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del query_validator, qp, train_idx, searcher\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG >= 1:\n    def display_stats(int_arr):\n        _arr = np.array(int_arr)\n        mean = np.mean(_arr)\n        median = np.median(_arr)\n        mode = stats.mode(_arr)\n        std_dev = np.std(_arr)\n        lprint(f'mean {mean}')\n        lprint(f'median {median}')\n        lprint(f'mode {mode}')\n        lprint(f'std dev {std_dev}')\n    \n    lprint('solution cover')\n    display_stats(curr_cover_dist)\n    lprint('max cover')\n    display_stats(max_cover_dist)\n    lprint('apk accuracy')\n    display_stats(apk_accuracy)\n    lprint('token efficiency stats')\n    display_stats(token_efficiency)\n    lprint('token usage stats')\n    display_stats(token_usage)\n    lprint('query time')\n    display_stats(query_time)\n    lprint('max time')\n    lprint(max_time)\n    lprint('total time')\n    lprint(total_time/60.0)\n    \n    final_score = np.mean(apks)\n\n    lprint(final_score)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pl.scan_csv(test_csv).select('publication_number').collect()\nsub = sub.with_columns(pl.Series(queries).alias('query'))\n\nsub","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.write_csv('submission.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}