{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59575,"databundleVersionId":8060720,"sourceType":"competition"},{"sourceId":8246447,"sourceType":"datasetVersion","datasetId":4892374},{"sourceId":8479599,"sourceType":"datasetVersion","datasetId":4517815},{"sourceId":8914164,"sourceType":"datasetVersion","datasetId":5360501},{"sourceId":9995420,"sourceType":"datasetVersion","datasetId":6151986},{"sourceId":174185912,"sourceType":"kernelVersion"}],"dockerImageVersionId":30786,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport pickle\nimport gc\nimport os\nimport random\nimport math\nimport copy\nfrom collections import defaultdict\nimport cupy as cp\n\n# Import custom utility functions (assuming whoosh_utils is provided)\nimport whoosh_utils\n\n# Install Whoosh library for indexing and searching text data\n# Uncomment the line below if Whoosh is not already installed\n# !pip install Whoosh==2.7.4\n\n# Configuration settings\nIS_TRAIN = False\nNEG_MAX_COUNT = 1000\npatents_per_word_patern = 3\nmagic = False\nPATENT_MAX_COUNT = 50 + NEG_MAX_COUNT  # Total allowed patents per query\nSUB_MAX_COUNT = 50 + NEG_MAX_COUNT  # Max number of patents in sub-search\nPATTERN_NUM_MAX = 10000\nCHAR_LIMIT = 9000  # Maximum character limit for queries\nNEG_WEIGHT = 0.1  # Weight for negative samples in scoring\nT0 = 1  # Initial temperature for simulated annealing\nT1 = 0.1  # Final temperature for simulated annealing\nmax_time = 1  # Maximum time for simulated annealing per sample in seconds\n\n# Paths for data\nTRAIN_PKL_PATH = '/kaggle/input/mappings-upsto'  # Update with your path\nfiltered_mappings_path = 'filtered_mappings'  # Directory to store filtered mappings\n\n# Create necessary directories\nos.makedirs(filtered_mappings_path, exist_ok=True)\n\n# Load nearest neighbors data\nif IS_TRAIN:\n    nn_path = '/kaggle/input/uspto-explainable-ai-validation-index/validation/validation_publication_numbers.csv'  # Update with your path\n    nn_df = pd.read_csv(nn_path)\n    print('Number of samples:', len(nn_df))\n    nn_df.to_csv('nn_df_for_index.csv', index=False)\nelse:\n    nn_df = pd.read_csv('/kaggle/input/uspto-explainable-ai/test.csv')  # Test data path\n    nn_df.to_csv('nn_df_for_index.csv', index=False)\n\n# Ensure necessary directories exist\nos.makedirs('reduce', exist_ok=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T12:48:28.819408Z","iopub.execute_input":"2024-11-25T12:48:28.820217Z","iopub.status.idle":"2024-11-25T12:49:11.606559Z","shell.execute_reply.started":"2024-11-25T12:48:28.820178Z","shell.execute_reply":"2024-11-25T12:49:11.605587Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Preprocessing: Filtering Patent-to-Word Mappings\n","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nfrom tqdm import tqdm\nimport pickle\nimport gc\n\n# Load the nearest neighbor data\nnn_df = pd.read_csv('nn_df_for_index.csv')\n\n# Collect all unique patents used in nn_df\nuse_patent = set()\nfor nn in nn_df.values:\n    for n in nn:\n        use_patent.add(n)\n\nprint('Number of unique patents used:', len(use_patent))\n\n# Initialize an empty dictionary to hold the filtered mappings\npatent_to_word_id = {}\n\n# Load and filter the patent_to_word_id mappings from chunks\nfor i in tqdm(range(21)):\n    chunk_path = f'/kaggle/input/mappings-upsto/patent_to_word_id_{i}.pkl'\n    with open(chunk_path, 'rb') as f:\n        _patent_to_word_id = pickle.load(f)\n    for patent, words in _patent_to_word_id.items():\n        if patent in use_patent:\n            patent_to_word_id[patent] = words\n    del _patent_to_word_id\n    gc.collect()\n\nprint('Number of patents after filtering:', len(patent_to_word_id))\n\n# Save the filtered mappings\nwith open(f'{filtered_mappings_path}patent_to_word_id.pkl', 'wb') as f:\n    pickle.dump(patent_to_word_id, f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T12:49:11.608372Z","iopub.execute_input":"2024-11-25T12:49:11.608678Z","iopub.status.idle":"2024-11-25T12:54:22.889295Z","shell.execute_reply.started":"2024-11-25T12:49:11.608650Z","shell.execute_reply":"2024-11-25T12:54:22.888433Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Mapping Patents to Numeric IDs","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport pickle\n\n# Load patent metadata to get all publication numbers\npatent_metadata_df = pd.read_parquet('/kaggle/input/uspto-explainable-ai/patent_metadata.parquet',\n                                     columns=['publication_number'])\n\n# Create mappings between publication numbers and numeric IDs\npatent_to_id = dict()\nid_to_patent = dict()\nfor idx, pub_num in enumerate(patent_metadata_df['publication_number'].values):\n    patent_to_id[pub_num] = idx\n    id_to_patent[idx] = pub_num\n\n# Save the mappings\nwith open(f'{filtered_mappings_path}patent_to_id.pkl', 'wb') as f:\n    pickle.dump(patent_to_id, f)\nwith open(f'{filtered_mappings_path}id_to_patent.pkl', 'wb') as f:\n    pickle.dump(id_to_patent, f)\nprint(\"saved patent to id\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T12:54:22.890448Z","iopub.execute_input":"2024-11-25T12:54:22.891433Z","iopub.status.idle":"2024-11-25T12:54:48.937691Z","shell.execute_reply.started":"2024-11-25T12:54:22.891392Z","shell.execute_reply":"2024-11-25T12:54:48.936849Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Filtering Word-to-Patent Mappings","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nfrom tqdm import tqdm\nimport pickle\nimport gc\nimport numpy as np\n\n# Load the mappings\npatent_to_id = pickle.load(open(f'{filtered_mappings_path}patent_to_id.pkl', 'rb'))\nid_to_patent = pickle.load(open(f'{filtered_mappings_path}id_to_patent.pkl', 'rb'))\nnn_df = pd.read_csv('nn_df_for_index.csv')\n\n# Convert publication numbers to numeric IDs in nn_df\nuse_patent_ids = set()\nfor nn in nn_df.values:\n    for n in nn:\n        use_patent_ids.add(patent_to_id[n])\n\nprint('Number of unique patent IDs used:', len(use_patent_ids))\n\n# Load the filtered patent_to_word_id\npatent_to_word_id = pickle.load(open(f'{filtered_mappings_path}patent_to_word_id.pkl', 'rb'))\n\n# Convert publication numbers to numeric IDs in patent_to_word_id\npatent_to_word_id_new = defaultdict(list)\nfor pub_num, words in tqdm(patent_to_word_id.items()):\n    patent_to_word_id_new[patent_to_id[pub_num]] = words\npatent_to_word_id = patent_to_word_id_new\n\n# Load word_to_patent_count\nbase_path = '/kaggle/input/mappings-upsto/'\nword_to_patent_count = pickle.load(open(base_path + 'word_to_patent_count.pkl', 'rb'))\n\n# Collect all words used\nuse_words = set()\nfor patent_id in tqdm(use_patent_ids):\n    words = patent_to_word_id[patent_id]\n    use_words.update(words)\nprint('Number of unique words used:', len(use_words))\n\n# Filter and save word_id_to_patent mappings\nfor i in tqdm(range(21)):\n    word_id_to_patent = {}\n    chunk_path = f'/kaggle/input/mappings-upsto/word_id_to_patent_{i}.pkl'\n    with open(chunk_path, 'rb') as f:\n        _word_id_to_patent = pickle.load(f)\n    for word_id, patent_set in _word_id_to_patent.items():\n        if word_id in use_words:\n            # Convert publication numbers to IDs\n            patent_set = np.array([patent_to_id[pub] for pub in patent_set], dtype='int32')\n            word_id_to_patent[word_id] = patent_set\n    with open(f'{filtered_mappings_path}word_id_to_patent_{i}.pkl', 'wb') as f:\n        pickle.dump(word_id_to_patent, f)\n    del word_id_to_patent, _word_id_to_patent\n    gc.collect()\n\ngc.collect()\n\n# Filter word_to_id and id_to_word mappings\nword_to_id = pickle.load(open(base_path + 'word_to_id.pkl', 'rb'))\nid_to_word = pickle.load(open(base_path + 'id_to_word.pkl', 'rb'))\n\n# Remove unused words\nremove_words = [w for w in id_to_word.keys() if w not in use_words]\nprint('Number of words to remove:', len(remove_words))\nfor word_id in remove_words:\n    word = id_to_word[word_id]\n    word_to_id.pop(word)\n    id_to_word.pop(word_id)\n\nprint('Number of words after filtering:', len(word_to_id))\n\n# Remove unused words from word_to_patent_count\nremove_words = [w for w in word_to_patent_count.keys() if w not in use_words]\nprint('Number of words to remove from word_to_patent_count:', len(remove_words))\nfor word in remove_words:\n    word_to_patent_count.pop(word)\nprint('Number of words in word_to_patent_count after filtering:', len(word_to_patent_count))\n\n# Save the filtered mappings\nwith open(f'{filtered_mappings_path}word_to_id.pkl', 'wb') as f:\n    pickle.dump(word_to_id, f)\nwith open(f'{filtered_mappings_path}id_to_word.pkl', 'wb') as f:\n    pickle.dump(id_to_word, f)\nwith open(f'{filtered_mappings_path}word_to_patent_count.pkl', 'wb') as f:\n    pickle.dump(word_to_patent_count, f)\nwith open(f'{filtered_mappings_path}patent_to_word_id.pkl', 'wb') as f:\n    pickle.dump(patent_to_word_id, f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T12:54:48.939704Z","iopub.execute_input":"2024-11-25T12:54:48.939977Z","iopub.status.idle":"2024-11-25T13:26:00.309488Z","shell.execute_reply.started":"2024-11-25T12:54:48.939951Z","shell.execute_reply":"2024-11-25T13:26:00.308587Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loading the Filtered Index","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nimport cupy as cp\n\n# Load word_id_to_patent mappings\nword_id_to_patent = {}\nfor i in tqdm(range(21)):\n    with open(f'{filtered_mappings_path}word_id_to_patent_{i}.pkl', 'rb') as f:\n        _word_id_to_patent = pickle.load(f)\n    for k, v in _word_id_to_patent.items():\n        # Convert to CuPy array for GPU acceleration\n        word_id_to_patent[k] = cp.array(np.sort(v), dtype=cp.int32)\n    del _word_id_to_patent\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:26:00.310745Z","iopub.execute_input":"2024-11-25T13:26:00.311017Z","iopub.status.idle":"2024-11-25T13:26:26.804585Z","shell.execute_reply.started":"2024-11-25T13:26:00.310991Z","shell.execute_reply":"2024-11-25T13:26:26.803743Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loading Additional Mappings","metadata":{}},{"cell_type":"code","source":"import pickle\n\n# Load word-to-ID and ID-to-word mappings\nbase_path = filtered_mappings_path\nword_to_id = pickle.load(open(base_path + 'word_to_id.pkl', 'rb'))\nid_to_word = pickle.load(open(base_path + 'id_to_word.pkl', 'rb'))\n\n# Load word-to-patent-count mapping\nword_to_patent_count = pickle.load(open(base_path + 'word_to_patent_count.pkl', 'rb'))\n\n# Load patent-to-word mapping\npatent_to_word_id = pickle.load(open(base_path + 'patent_to_word_id.pkl', 'rb'))\n\n# Load patent ID mappings\npatent_to_id = pickle.load(open(base_path + 'patent_to_id.pkl', 'rb'))\nid_to_patent = pickle.load(open(base_path + 'id_to_patent.pkl', 'rb'))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:26:26.805887Z","iopub.execute_input":"2024-11-25T13:26:26.806243Z","iopub.status.idle":"2024-11-25T13:26:38.131197Z","shell.execute_reply.started":"2024-11-25T13:26:26.806205Z","shell.execute_reply":"2024-11-25T13:26:38.130477Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyzing Word Frequencies","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.ticker as ticker\nfrom itertools import chain\n\n# Assuming word_to_patent_count is already defined and populated\ncounts = list(word_to_patent_count.values())\n\nplt.figure(figsize=(12, 7))\nplt.hist(counts, bins=20, log=True, color='skyblue', edgecolor='black')\n\nplt.title('Distribution of Word Frequencies (Log Scale)', fontsize=18)\nplt.xlabel('Number of Patents Containing the Word', fontsize=14)\nplt.ylabel('Frequency (Log Scale)', fontsize=14)\n\nax = plt.gca()\nax.set_yscale('log')\n\n# Major ticks\nax.yaxis.set_major_locator(ticker.LogLocator(base=10.0, numticks=15))\nax.yaxis.set_major_formatter(ticker.FuncFormatter(lambda y, _: f'{y:g}'))\n\n# Minor ticks\nax.yaxis.set_minor_locator(ticker.LogLocator(base=10.0, subs='auto', numticks=15))\nax.yaxis.set_minor_formatter(ticker.NullFormatter())\n\n# Customize tick label font size\nax.tick_params(axis='both', which='major', labelsize=12)\n\n# Enhance grid\nplt.grid(True, which=\"both\", ls=\"--\", linewidth=0.5)\nplt.grid(True, which=\"minor\", ls=\":\", linewidth=0.3)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:26:38.132310Z","iopub.execute_input":"2024-11-25T13:26:38.132550Z","iopub.status.idle":"2024-11-25T13:26:39.182096Z","shell.execute_reply.started":"2024-11-25T13:26:38.132526Z","shell.execute_reply":"2024-11-25T13:26:39.181209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Preparing Nearest Neighbors Data","metadata":{}},{"cell_type":"code","source":"# Convert publication numbers to numeric IDs in nn_df\nfor i in range(nn_df.values.shape[0]):\n    for j in range(nn_df.values.shape[1]):\n        nn_df.values[i, j] = patent_to_id[nn_df.values[i, j]]\n\n# Extract neighbor IDs (excluding the query patent)\nneighbors = nn_df.values[:, 1:]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:26:39.183431Z","iopub.execute_input":"2024-11-25T13:26:39.184130Z","iopub.status.idle":"2024-11-25T13:26:39.191485Z","shell.execute_reply.started":"2024-11-25T13:26:39.184088Z","shell.execute_reply":"2024-11-25T13:26:39.190658Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Generating Word Patterns","metadata":{}},{"cell_type":"code","source":"import itertools\nimport time\n\nword_pattern_to_patent_set_list = [defaultdict(set) for _ in range(len(neighbors))]\nword_pattern_to_patent_set_all = dict()\nadd_count_list = []\n\nfor i in tqdm(range(len(neighbors))):\n    word_pattern_to_patent_set = word_pattern_to_patent_set_list[i]\n    add_list = []\n    nn_list = neighbors[i]\n    nn_array = np.sort(nn_list)\n    nn_array = cp.array(nn_array, dtype=cp.int32)\n    nn_set = set(nn_list)\n    use_words_cache = {}\n    t_sum = 0\n\n    for r in range(1, 3):  # Considering combinations\n        debug_count = 0\n        for nn_comb in itertools.combinations(nn_list, r):\n            counts = []\n            if nn_comb[0] not in use_words_cache:\n                use_words_cache[nn_comb[0]] = set(patent_to_word_id[nn_comb[0]])\n            use_words = use_words_cache[nn_comb[0]].copy()\n            for nn in nn_comb[1:]:\n                if nn not in use_words_cache:\n                    use_words_cache[nn] = set(patent_to_word_id[nn])\n                use_words &= use_words_cache[nn]\n            use_words = tuple(sorted(use_words))\n\n            if len(use_words) == 0:\n                continue\n\n            use_words_sorted = sorted(use_words, key=lambda x: word_to_patent_count[x])\n            use_words = tuple([w for w in use_words if w in set(use_words_sorted)])\n            use_words_final = []\n\n            # Initialize with the first word\n            all_patent_set = word_id_to_patent[use_words_sorted[0]]\n            use_words_final.append(use_words_sorted[0])\n            cpc_count = 1 if id_to_word[use_words_sorted[0]].startswith('cpc') else 0\n\n            # Intersect with other words\n            if not len(all_patent_set) <= r:\n                for word in use_words_sorted[1:]:\n                    if id_to_word[word].startswith('cpc') and cpc_count == 1:\n                        continue\n\n                    before_len = len(all_patent_set)\n                    all_patent_set = cp.intersect1d(all_patent_set, word_id_to_patent[word], assume_unique=True)\n                    after_len = len(all_patent_set)\n\n                    if after_len < before_len:\n                        use_words_final.append(word)\n                        if id_to_word[word].startswith('cpc'):\n                            cpc_count += 1\n\n                    if len(all_patent_set) <= r:\n                        break\n\n                    cnt1 = len(all_patent_set)\n                    if cnt1 <= 50:\n                        t1 = time.time()\n                        cnt2 = len(cp.intersect1d(all_patent_set, nn_array, assume_unique=True))\n                        t_sum += time.time() - t1\n                        if cnt1 == cnt2:\n                            break\n\n            use_words = tuple(use_words_final)\n            if len(all_patent_set) <= 50:\n                all_patent_set = cp.asnumpy(all_patent_set)\n                nn_comb_set = set(nn_comb)\n                all_patent_set = set(all_patent_set)\n                add_set = all_patent_set & nn_set\n                nn_comb_set = nn_comb_set | add_set\n                if len(nn_comb_set) == len(all_patent_set):\n                    add_list.append((use_words, nn_comb_set, all_patent_set))\n\n        if i <= 50:\n            print(f'Time for r={r}:', t_sum)\n            print(f'r >= {r} len(add_list): {len(add_list)}')\n            print('Max use_words length:', len(use_words_final))\n\n    add_count = 0\n    for word_pattern, patent_set, all_patent_set in add_list:\n        if word_pattern in word_pattern_to_patent_set:\n            continue\n        word_pattern_to_patent_set[word_pattern] = patent_set\n        word_pattern_to_patent_set_all[word_pattern] = all_patent_set\n        add_count += 1\n\n    if i <= 50:\n        print(f'len(add_list): {len(add_list)}, add_count: {add_count}')\n    add_count_list.append(add_count)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:26:39.193069Z","iopub.execute_input":"2024-11-25T13:26:39.193508Z","iopub.status.idle":"2024-11-25T13:28:46.681218Z","shell.execute_reply.started":"2024-11-25T13:26:39.193467Z","shell.execute_reply":"2024-11-25T13:28:46.680314Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Filtering Word Patterns","metadata":{}},{"cell_type":"code","source":"def count_query_len(word_pattern):\n    # For simplicity, we consider a fixed token length\n    token_len = 2\n    return token_len\n\ndef filtering_word_pattern_to_patent_set_list(word_pattern_to_patent_set_list):\n    word_pattern_to_patent_set_list_new = []\n    for word_pattern_to_patent_set in word_pattern_to_patent_set_list:\n        word_pattern_to_patent_set_new = defaultdict(set)\n        sort_keys = []\n        for word_pattern, patent_set in word_pattern_to_patent_set.items():\n            max_count = len(word_pattern_to_patent_set_all[word_pattern])\n            if max_count <= PATENT_MAX_COUNT:\n                neg_patent_set = word_pattern_to_patent_set_all[word_pattern] - patent_set\n                neg_weight = len(neg_patent_set)\n                sort_key = (len(patent_set), -count_query_len(word_pattern), -neg_weight)\n                sort_keys.append((sort_key, word_pattern))\n                word_pattern_to_patent_set_new[word_pattern] = patent_set\n\n        # Sort patterns based on the sort keys\n        sorted_word_patterns = sorted(sort_keys, key=lambda x: x[0], reverse=True)\n        word_pattern_to_patent_set_new = [(word_pattern, word_pattern_to_patent_set_new[word_pattern]) for _, word_pattern in sorted_word_patterns]\n\n        # Remove duplicates\n        seen = set()\n        new_state = []\n        for word_pattern, patent_set in word_pattern_to_patent_set_new:\n            pattern = (tuple(sorted(list(patent_set))), count_query_len(word_pattern))\n            if pattern in seen:\n                continue\n            seen.add(pattern)\n            new_state.append((word_pattern, patent_set))\n        word_pattern_to_patent_set_new = new_state\n\n        # Limit the number of patterns\n        word_pattern_to_patent_set_new = word_pattern_to_patent_set_new[:PATTERN_NUM_MAX]\n        word_pattern_to_patent_set_list_new.append(word_pattern_to_patent_set_new)\n\n    return word_pattern_to_patent_set_list_new\n\nword_pattern_to_patent_set_list = filtering_word_pattern_to_patent_set_list(word_pattern_to_patent_set_list)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:28:46.684022Z","iopub.execute_input":"2024-11-25T13:28:46.684360Z","iopub.status.idle":"2024-11-25T13:28:46.810322Z","shell.execute_reply.started":"2024-11-25T13:28:46.684332Z","shell.execute_reply":"2024-11-25T13:28:46.809643Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyzing Recall Distribution","metadata":{}},{"cell_type":"code","source":"# Calculate recall for each sample\nrecall_list = []\nfor word_pattern_to_patent_set in word_pattern_to_patent_set_list:\n    true_set = set()\n    for word_pattern, true_patent_set in word_pattern_to_patent_set:\n        true_set |= true_patent_set\n    recall_list.append(len(true_set))\n\nprint(pd.DataFrame(recall_list).describe())\nplt.hist(recall_list)\nplt.title('Recall Distribution')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:46:27.410268Z","iopub.execute_input":"2024-11-25T13:46:27.411176Z","iopub.status.idle":"2024-11-25T13:46:27.666054Z","shell.execute_reply.started":"2024-11-25T13:46:27.411128Z","shell.execute_reply":"2024-11-25T13:46:27.665143Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Simulated Annealing for Query Optimization","metadata":{}},{"cell_type":"code","source":"from functools import total_ordering\n\n@total_ordering\nclass State:\n    def __init__(self, use_word_pattern_list, not_use_word_pattern_list, use_true_patent_set, neg_count, score=None,\n                 is_hard_penalty=False):\n        self.use_word_pattern_list = use_word_pattern_list\n        self.not_use_word_pattern_list = not_use_word_pattern_list\n        self.use_true_patent_set = use_true_patent_set\n        self.neg_count = neg_count\n        self.is_hard_penalty = is_hard_penalty\n\n    def __lt__(self, other):\n        return self.score < other.score\n\n    def __eq__(self, other):\n        return self.score == other.score\n\n    def calc_score(self):\n        if self.is_hard_penalty:\n            self.score = len(self.use_true_patent_set) - NEG_WEIGHT * self.neg_count\n        else:\n            self.score = len(self.use_true_patent_set) - NEG_WEIGHT * self.neg_count\n\nclass Timer:\n    def __init__(self):\n        self.start = time.time()\n\n    def get_current_time(self):\n        return (time.time() - self.start)\n\n# Calculate acceptance probability for simulated annealing\ndef calc_sa_p(new_score, score, T):\n    score_diff = new_score - score\n    if score_diff >= 0:\n        return 1\n    else:\n        return math.exp(score_diff / T)\n\n# Average precision at 50\ndef ap50(preds, labels):\n    precisions = list()\n    n_label = len(labels)\n    n_found = 0\n    for e, i in enumerate(preds):\n        if i in labels:\n            n_found += 1\n        precisions.append(n_found/(e+1))\n    return sum(precisions)/50\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:48:40.896922Z","iopub.execute_input":"2024-11-25T13:48:40.897279Z","iopub.status.idle":"2024-11-25T13:48:40.905685Z","shell.execute_reply.started":"2024-11-25T13:48:40.897246Z","shell.execute_reply":"2024-11-25T13:48:40.904773Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Running Simulated Annealing","metadata":{}},{"cell_type":"code","source":"# Prepare necessary data structures\nword_pattern_to_neg_list = []\n\nfor sample_idx, word_pattern_to_patent_set in tqdm(enumerate(word_pattern_to_patent_set_list), total=len(word_pattern_to_patent_set_list)):\n    word_pattern_to_neg_count = dict()\n    for c, patent_set in word_pattern_to_patent_set:\n        neg_patent_set = word_pattern_to_patent_set_all[c] - patent_set\n        word_pattern_to_neg_count[c] = len(neg_patent_set)\n    word_pattern_to_neg_list.append(word_pattern_to_neg_count)\n\ndef word_pattern_to_query(word_pattern, magic=True):\n    query_rm_cpc = [s for s in word_pattern if not id_to_word[s].startswith('cpc')]\n    query_cpc = [s for s in word_pattern if id_to_word[s].startswith('cpc')]\n    word_pattern = query_rm_cpc + query_cpc\n    query = ''\n    for i, word in enumerate(word_pattern):\n        word_str = id_to_word[word]\n        if word_str.startswith('cpc'):\n            if i != len(word_pattern) - 1:\n                query += f'{word_str}*'\n                base, suffix = word_str.split('/')\n                suffix = suffix[:-1] + '?'\n                query += f' {base}/{suffix}'\n            else:\n                query += f'{word_str}'\n        else:\n            if magic:\n                query += f'{word_str}-'\n            else:\n                query += f'{word_str} AND '\n    if magic:\n        if query.endswith('-'):\n            query = query[:-1]\n    else:\n        if query.endswith(' AND '):\n            query = query[:-5]\n    query = \"(\" + query + ')'\n    return query\n\n\n# Precompute character lengths of queries\nword_pattern_to_char_len = dict()\nfor word_pattern_to_patent_set in word_pattern_to_patent_set_list:\n    for word_pattern, _ in word_pattern_to_patent_set:\n        query = word_pattern_to_query(word_pattern, magic=magic)\n        word_pattern_to_char_len[word_pattern] = len(query) + 4\n\n# Initialize variables for results\nquery_list = []\ndefault_query = 'ti:titonium'\nsuccess_count = 0\nscore_list = []\ntrue_patent_set_len_list = [0] * len(word_pattern_to_patent_set_list)\npatent_set_len_list = [0] * len(word_pattern_to_patent_set_list)\nneg_count_list = [0] * len(word_pattern_to_patent_set_list)\n\nfor sample_idx, word_pattern_to_patent_set in tqdm(enumerate(word_pattern_to_patent_set_list), total=len(word_pattern_to_patent_set_list)):\n    word_pattern_to_neg_count = word_pattern_to_neg_list[sample_idx]\n    true_set = set(nn_df.values[sample_idx, 1:])\n    labels = list(nn_df.values[sample_idx, 1:])\n    labels = [id_to_patent[patent_id] for patent_id in labels]\n\n    # Handle cases with no word patterns\n    if len(word_pattern_to_patent_set) == 0:\n        score_list.append(0)\n        query_list.append(default_query)\n        continue\n\n    word_pattern_to_patent_set_dict = dict()\n    for word_pattern, patent_set in word_pattern_to_patent_set:\n        word_pattern_to_patent_set_dict[word_pattern] = patent_set\n\n    use_word_pattern_list = []\n    not_use_word_pattern_list = []\n    for word_pattern, _ in word_pattern_to_patent_set:\n        not_use_word_pattern_list.append(word_pattern)\n\n    # Precompute query lengths\n    word_pattern_to_query_len = dict()\n    for word_pattern, _ in word_pattern_to_patent_set:\n        word_pattern_to_query_len[word_pattern] = count_query_len(word_pattern)\n\n    curr_state = State([], [], set(), 0, 0)\n    curr_state.score = 0\n    curr_query_len = -1\n    best_query_len = -1\n    curr_char_len = -4\n    best_char_len = -4\n    best_state = copy.deepcopy(curr_state)\n    best_state.score = 0\n    timer = Timer()\n    _score_list = []\n    pattern_count = [0, 0]\n    is_hard_penalty = False\n\n    while True:\n        curr_time = timer.get_current_time()\n        if curr_time > max_time:\n            break\n        t = curr_time / max_time\n        T = T0**(1-t) * T1**t\n        p = random.random()\n        act = 'add_pattern' if p >= 0.5 else 'remove_pattern'\n        next_state = State([], [], set(), 0, 0, is_hard_penalty)\n\n        if act == 'add_pattern':\n            N = len(not_use_word_pattern_list)\n            if N == 0:\n                continue\n            idx = random.randint(0, N-1)\n            c = not_use_word_pattern_list[idx]\n            query_len = word_pattern_to_query_len[c]\n            char_len = word_pattern_to_char_len[c]\n            if curr_query_len + query_len > 50 or curr_char_len + char_len > CHAR_LIMIT:\n                is_exceed_query_limit = True\n            else:\n                is_exceed_query_limit = False\n        elif act == 'remove_pattern':\n            N = len(use_word_pattern_list)\n            if N == 0:\n                continue\n            idx = random.randint(0, N-1)\n            c = use_word_pattern_list[idx]\n            query_len = word_pattern_to_query_len[c]\n            char_len = word_pattern_to_char_len[c]\n            is_exceed_query_limit = False\n\n        if is_exceed_query_limit:\n            continue\n\n        # Update state\n        if act == 'add_pattern':\n            neg_count = curr_state.neg_count + word_pattern_to_neg_count[c]\n            use_true_patent_set = curr_state.use_true_patent_set | word_pattern_to_patent_set_dict[c]\n            update_count = len(use_true_patent_set) - len(curr_state.use_true_patent_set)\n            if update_count == 0:\n                continue\n        else:\n            neg_count = 0\n            use_true_patent_set = set()\n            for word_pattern in use_word_pattern_list:\n                if act == 'remove_pattern' and word_pattern == c:\n                    continue\n                neg_count += word_pattern_to_neg_count[word_pattern]\n                if len(use_true_patent_set) + neg_count > SUB_MAX_COUNT:\n                    break\n                use_true_patent_set |= word_pattern_to_patent_set_dict[word_pattern]\n\n        if len(use_true_patent_set) + neg_count > SUB_MAX_COUNT or neg_count > NEG_MAX_COUNT:\n            continue\n\n        next_state.neg_count = neg_count\n        next_state.use_true_patent_set = use_true_patent_set\n        next_state.calc_score()\n\n        # Acceptance probability\n        sa_p = calc_sa_p(next_state.score, curr_state.score, T)\n        if random.random() < sa_p:\n            curr_state = next_state\n            if act == 'add_pattern':\n                c = not_use_word_pattern_list.pop(idx)\n                use_word_pattern_list.append(c)\n                curr_query_len += query_len\n                curr_char_len += char_len\n            elif act == 'remove_pattern':\n                c = use_word_pattern_list.pop(idx)\n                not_use_word_pattern_list.append(c)\n                curr_query_len -= query_len\n                curr_char_len -= char_len\n\n        _score_list.append(curr_state.score)\n\n        if curr_state.score > best_state.score:\n            best_state = curr_state\n            best_state.use_word_pattern_list = copy.deepcopy(use_word_pattern_list)\n            best_query_len = curr_query_len\n            best_char_len = curr_char_len\n            if best_state.score >= 35 and not is_hard_penalty:\n                is_hard_penalty = True\n                curr_state.is_hard_penalty = True\n                curr_state.calc_score()\n                best_state.is_hard_penalty = True\n                best_state.calc_score()\n\n        if act == 'add_pattern':\n            pattern_count[0] += 1\n        else:\n            pattern_count[1] += 1\n\n    # Generate the final query\n    if best_state is None:\n        print('Error: best_state is None')\n        query_list.append(default_query)\n    else:\n        use_true_patent_set = set()\n        use_patent_set = set()\n        for c in best_state.use_word_pattern_list:\n            use_patent_set |= word_pattern_to_patent_set_all[c]\n            use_true_patent_set |= word_pattern_to_patent_set_dict[c]\n        use_word_pattern_list = best_state.use_word_pattern_list\n        true_patent_set_len_list[sample_idx] = len(use_true_patent_set)\n        patent_set_len_list[sample_idx] = len(use_patent_set)\n        neg_count_list[sample_idx] = best_state.neg_count\n        query = ''\n        for i, word_pattern in enumerate(use_word_pattern_list):\n            prev_query = query\n            _query = word_pattern_to_query(word_pattern, magic=magic)\n            query += _query\n            if i != len(use_word_pattern_list) - 1:\n                query += ' OR '\n            if whoosh_utils.count_query_tokens(query) > 50 or len(query) > CHAR_LIMIT + 500:\n                print('Query limit exceeded')\n                query = prev_query\n                break\n        query_list.append(query)\n        success_count += 1\n        if IS_TRAIN:\n            results = whoosh_utils.execute_query(query, qp, searcher)\n            result_set = set(results)\n            n_pick = len(set(labels) & result_set)\n            score = ap50(results + [-1] * (50 - len(results)), labels)\n        else:\n            score = 0\n            n_pick = 0\n        if sample_idx < 50:\n            print('Query:', query)\n            print('AP@50:', score)\n            print('Number of correct picks:', n_pick)\n            print('True patent set size:', true_patent_set_len_list[sample_idx])\n            print('Negative count:', neg_count_list[sample_idx])\n            print('Query length:', len(query), 'Characters:', best_char_len)\n        score_list.append(score)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:46:34.360735Z","iopub.execute_input":"2024-11-25T13:46:34.361047Z","iopub.status.idle":"2024-11-25T13:46:44.614775Z","shell.execute_reply.started":"2024-11-25T13:46:34.361020Z","shell.execute_reply":"2024-11-25T13:46:44.613852Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyzing True Patent Set Length Distribution","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.ticker as ticker\n\n# Calculate basic statistics\nmean_true_patent_set_len = np.mean(true_patent_set_len_list)\nmedian_true_patent_set_len = np.median(true_patent_set_len_list)\nstd_true_patent_set_len = np.std(true_patent_set_len_list)\n\nprint(\"Statistics for true_patent_set_len_list:\")\nprint(f\"Mean: {mean_true_patent_set_len:.2f}\")\nprint(f\"Median: {median_true_patent_set_len:.2f}\")\nprint(f\"Standard Deviation: {std_true_patent_set_len:.2f}\")\n\n# Define the figure size for better readability\nplt.figure(figsize=(12, 7))\n\n# Plot the histogram with customized bins and aesthetics\nn, bins, patches = plt.hist(\n    true_patent_set_len_list, \n    bins=50,  # Increased number of bins for better resolution\n    color='skyblue', \n    edgecolor='black', \n    alpha=0.7,  # Added transparency\n    density=False  # Set to True if you want probability density instead of counts\n)\n\n# Add title and labels with increased font sizes\nplt.title('Distribution of True Patent Set Length', fontsize=16)\nplt.xlabel('Number of Patents', fontsize=14)\nplt.ylabel('Frequency', fontsize=14)\n\n# Add vertical lines for Mean and Median\nplt.axvline(mean_true_patent_set_len, color='red', linestyle='dashed', linewidth=2, label=f'Mean: {mean_true_patent_set_len:.2f}')\nplt.axvline(median_true_patent_set_len, color='green', linestyle='dashed', linewidth=2, label=f'Median: {median_true_patent_set_len:.2f}')\n\n# Add a legend to identify Mean and Median lines\nplt.legend(fontsize=12)\n\n# Customize y-axis to prevent scientific notation if necessary\nax = plt.gca()\nax.yaxis.set_major_formatter(ticker.ScalarFormatter())\nax.yaxis.get_major_formatter().set_scientific(False)\nax.yaxis.get_major_formatter().set_useOffset(False)\n\n# Optionally, set y-axis to logarithmic scale if data is highly skewed\n# Uncomment the following lines if needed\n# ax.set_yscale('log')\n# ax.yaxis.set_major_locator(ticker.LogLocator(base=10.0, numticks=15))\n# ax.yaxis.set_major_formatter(ticker.FuncFormatter(lambda y, _: f'{y:g}'))\n\n# Add grid lines for better readability\nplt.grid(True, which='both', linestyle='--', linewidth=0.5, alpha=0.7)\n\n# Adjust layout to prevent clipping of labels and titles\nplt.tight_layout()\n\n# Display the plot\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:47:22.830328Z","iopub.execute_input":"2024-11-25T13:47:22.830692Z","iopub.status.idle":"2024-11-25T13:47:23.318129Z","shell.execute_reply.started":"2024-11-25T13:47:22.830660Z","shell.execute_reply":"2024-11-25T13:47:23.317192Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyzing Score Distribution","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.ticker as ticker\nfrom itertools import chain\n\n# Calculate basic statistics\nmean_score = np.mean(score_list)\nmedian_score = np.median(score_list)\nstd_score = np.std(score_list)\n\nprint(\"Statistics for score_list:\")\nprint(f\"Mean: {mean_score:.2f}\")\nprint(f\"Median: {median_score:.2f}\")\nprint(f\"Standard Deviation: {std_score:.2f}\")\n\n# Check if all scores are identical\nall_scores_identical = all(score == score_list[0] for score in score_list)\n\n# Define the figure size for better readability\nplt.figure(figsize=(12, 7))\n\nif all_scores_identical:\n    # All scores are identical; use a single bar chart or text annotation\n    unique_score = score_list[0]\n    count = len(score_list)\n    \n    # Option 1: Single Bar Chart\n    plt.bar(unique_score, count, color='skyblue', edgecolor='black', width=0.5)\n    plt.title('Distribution of Scores', fontsize=16)\n    plt.xlabel('Score', fontsize=14)\n    plt.ylabel('Frequency', fontsize=14)\n    plt.xticks([unique_score])  # Single tick at the unique score\n    plt.text(unique_score, count, f'Count: {count}', ha='center', va='bottom', fontsize=12)\n    \n    # Option 2: Text Annotation (Uncomment if preferred)\n    # plt.text(0.5, 0.5, f'All Scores = {unique_score}\\nCount = {count}', \n    #          horizontalalignment='center', \n    #          verticalalignment='center', \n    #          fontsize=14, \n    #          bbox=dict(facecolor='skyblue', alpha=0.5, boxstyle='round,pad=1'))\n    # plt.title('Score Distribution')\n    # plt.axis('off')  # Hide the axes\nelse:\n    # Scores vary; plot histogram with enhancements\n    n, bins, patches = plt.hist(\n        score_list, \n        bins=50,  # Increased number of bins for better resolution\n        color='skyblue', \n        edgecolor='black', \n        alpha=0.7,  # Added transparency\n        density=False  # Set to True if you want probability density instead of counts\n    )\n    \n    # Add title and labels with increased font sizes\n    plt.title('Distribution of Scores', fontsize=16)\n    plt.xlabel('Score', fontsize=14)\n    plt.ylabel('Frequency', fontsize=14)\n    \n    # Add vertical lines for Mean and Median\n    plt.axvline(mean_score, color='red', linestyle='dashed', linewidth=2, label=f'Mean: {mean_score:.2f}')\n    plt.axvline(median_score, color='green', linestyle='dashed', linewidth=2, label=f'Median: {median_score:.2f}')\n    \n    # Add a legend to identify Mean and Median lines\n    plt.legend(fontsize=12)\n    \n    # Customize y-axis to prevent scientific notation and improve readability\n    ax = plt.gca()\n    ax.yaxis.set_major_formatter(ticker.ScalarFormatter())\n    ax.yaxis.get_major_formatter().set_scientific(False)\n    ax.yaxis.get_major_formatter().set_useOffset(False)\n    \n    # Optionally, set y-axis to logarithmic scale if data is highly skewed\n    # Uncomment the following lines if needed\n    # ax.set_yscale('log')\n    # ax.yaxis.set_major_locator(ticker.LogLocator(base=10.0, numticks=15))\n    # ax.yaxis.set_major_formatter(ticker.FuncFormatter(lambda y, _: f'{y:g}'))\n    \n    # Add grid lines for better readability\n    plt.grid(True, which='both', linestyle='--', linewidth=0.5, alpha=0.7)\n    \n    # Adjust layout to prevent clipping of labels and titles\n    plt.tight_layout()\n\n# Display the plot\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:47:34.972005Z","iopub.execute_input":"2024-11-25T13:47:34.972354Z","iopub.status.idle":"2024-11-25T13:47:35.191996Z","shell.execute_reply.started":"2024-11-25T13:47:34.972322Z","shell.execute_reply":"2024-11-25T13:47:35.191053Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyzing Patent Set Length Distribution","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.ticker as ticker\nfrom itertools import chain\n\n# Calculate basic statistics\nmean_patent_set_len = np.mean(patent_set_len_list)\nmedian_patent_set_len = np.median(patent_set_len_list)\nstd_patent_set_len = np.std(patent_set_len_list)\n\nprint(\"Statistics for patent_set_len_list:\")\nprint(f\"Mean: {mean_patent_set_len:.2f}\")\nprint(f\"Median: {median_patent_set_len:.2f}\")\nprint(f\"Standard Deviation: {std_patent_set_len:.2f}\")\n\n# Check if all patent_set_len_list values are identical\nall_patent_set_len_identical = all(patent_len == patent_set_len_list[0] for patent_len in patent_set_len_list)\n\n# Define the figure size for better readability\nplt.figure(figsize=(12, 7))\n\nif all_patent_set_len_identical:\n    # All patent_set_len_list values are identical; use a single bar chart or text annotation\n    unique_patent_len = patent_set_len_list[0]\n    count = len(patent_set_len_list)\n    \n    # Option 1: Single Bar Chart\n    plt.bar(unique_patent_len, count, color='salmon', edgecolor='black', width=0.5)\n    plt.title('Distribution of Patent Set Length', fontsize=16)\n    plt.xlabel('Number of Patents', fontsize=14)\n    plt.ylabel('Frequency', fontsize=14)\n    plt.xticks([unique_patent_len])  # Single tick at the unique value\n    plt.text(unique_patent_len, count, f'Count: {count}', ha='center', va='bottom', fontsize=12)\n    \n    # Option 2: Text Annotation (Uncomment if preferred)\n    # plt.text(0.5, 0.5, f'All Patent Set Lengths = {unique_patent_len}\\nCount = {count}', \n    #          horizontalalignment='center', \n    #          verticalalignment='center', \n    #          fontsize=14, \n    #          bbox=dict(facecolor='salmon', alpha=0.5, boxstyle='round,pad=1'))\n    # plt.title('Patent Set Length Distribution')\n    # plt.axis('off')  # Hide the axes\nelse:\n    # Patent set lengths vary; plot histogram with enhancements\n    n, bins, patches = plt.hist(\n        patent_set_len_list, \n        bins=50,  # Increased number of bins for better resolution\n        color='salmon', \n        edgecolor='black', \n        alpha=0.7,  # Added transparency\n        density=False  # Set to True if you want probability density instead of counts\n    )\n    \n    # Add title and labels with increased font sizes\n    plt.title('Distribution of Patent Set Lengths', fontsize=16)\n    plt.xlabel('Number of Patents', fontsize=14)\n    plt.ylabel('Frequency', fontsize=14)\n    \n    # Add vertical lines for Mean and Median\n    plt.axvline(mean_patent_set_len, color='blue', linestyle='dashed', linewidth=2, label=f'Mean: {mean_patent_set_len:.2f}')\n    plt.axvline(median_patent_set_len, color='green', linestyle='dashed', linewidth=2, label=f'Median: {median_patent_set_len:.2f}')\n    \n    # Add a legend to identify Mean and Median lines\n    plt.legend(fontsize=12)\n    \n    # Customize y-axis to prevent scientific notation and improve readability\n    ax = plt.gca()\n    ax.yaxis.set_major_formatter(ticker.ScalarFormatter())\n    ax.yaxis.get_major_formatter().set_scientific(False)\n    ax.yaxis.get_major_formatter().set_useOffset(False)\n    \n    # Optionally, set y-axis to logarithmic scale if data is highly skewed\n    # Uncomment the following lines if needed\n    # ax.set_yscale('log')\n    # ax.yaxis.set_major_locator(ticker.LogLocator(base=10.0, numticks=15))\n    # ax.yaxis.set_major_formatter(ticker.FuncFormatter(lambda y, _: f'{y:g}'))\n    \n    # Add grid lines for better readability\n    plt.grid(True, which='both', linestyle='--', linewidth=0.5, alpha=0.7)\n    \n    # Adjust layout to prevent clipping of labels and titles\n    plt.tight_layout()\n\n# Display the plot\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:48:51.562176Z","iopub.execute_input":"2024-11-25T13:48:51.562842Z","iopub.status.idle":"2024-11-25T13:48:52.028333Z","shell.execute_reply.started":"2024-11-25T13:48:51.562805Z","shell.execute_reply":"2024-11-25T13:48:52.027517Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyzing Negative Count Distribution","metadata":{}},{"cell_type":"code","source":"# Calculate the mean of neg_count_list\nmean_neg_count = np.mean(neg_count_list)\nprint(\"Mean of neg_count_list:\", mean_neg_count)\n\n# Plot histogram\nplt.hist(neg_count_list)\nplt.title('neg_count_list distribution')\nplt.xlabel('Negative Count')\nplt.ylabel('Frequency')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:48:55.899645Z","iopub.execute_input":"2024-11-25T13:48:55.900386Z","iopub.status.idle":"2024-11-25T13:48:56.065198Z","shell.execute_reply.started":"2024-11-25T13:48:55.900354Z","shell.execute_reply":"2024-11-25T13:48:56.063843Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Computing Correct Negative Counts","metadata":{}},{"cell_type":"code","source":"# Compute correct neg_count_list based on total and true patent counts\nneg_count_list_2 = [patent_set_len_list[i] - true_patent_set_len_list[i] for i in range(len(patent_set_len_list))]\nmean_neg_count_2 = np.mean(neg_count_list_2)\nprint(\"Mean of neg_count_list_2:\", mean_neg_count_2)\n\n# Plot histogram\nplt.hist(neg_count_list_2)\nplt.title('neg_count_list_2 distribution')\nplt.xlabel('Correct Negative Count')\nplt.ylabel('Frequency')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:47:49.671585Z","iopub.execute_input":"2024-11-25T13:47:49.672436Z","iopub.status.idle":"2024-11-25T13:47:49.835268Z","shell.execute_reply.started":"2024-11-25T13:47:49.672400Z","shell.execute_reply":"2024-11-25T13:47:49.834058Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyzing Difference in Negative Counts","metadata":{}},{"cell_type":"code","source":"# Compute the difference between the two negative count lists\nneg_count_list_diff = [neg_count_list[i] - neg_count_list_2[i] for i in range(len(neg_count_list))]\nmean_neg_count_diff = np.mean(neg_count_list_diff)\nprint(\"Mean difference between neg_count_list and neg_count_list_2:\", mean_neg_count_diff)\n\n# Plot histogram\nplt.hist(neg_count_list_diff)\nplt.title('neg_count_list_diff distribution')\nplt.xlabel('Difference in Negative Counts')\nplt.ylabel('Frequency')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:48:12.737087Z","iopub.execute_input":"2024-11-25T13:48:12.737426Z","iopub.status.idle":"2024-11-25T13:48:12.974402Z","shell.execute_reply.started":"2024-11-25T13:48:12.737396Z","shell.execute_reply":"2024-11-25T13:48:12.973449Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Displaying Scores or Sample Scores","metadata":{}},{"cell_type":"code","source":"if IS_TRAIN:\n    print(\"First 50 values of true_patent_set_len_list:\")\n    print(true_patent_set_len_list[:50])\n    print(\"First 50 values of score_list:\")\n    print(score_list[:50])\nelse:\n    print(\"First 10 values of score_list:\")\n    print(score_list[:10])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:48:15.753621Z","iopub.execute_input":"2024-11-25T13:48:15.753981Z","iopub.status.idle":"2024-11-25T13:48:15.758887Z","shell.execute_reply.started":"2024-11-25T13:48:15.753950Z","shell.execute_reply":"2024-11-25T13:48:15.758057Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyzing Query Token Length Distribution","metadata":{}},{"cell_type":"code","source":"# Compute query token counts\nquery_count_list = [whoosh_utils.count_query_tokens(query) for query in query_list]\nmax_query_tokens = max(query_count_list)\nprint(\"Maximum query token count:\", max_query_tokens)\n\n# Plot histogram\nplt.hist(query_count_list)\nplt.title('Query Token Length Distribution')\nplt.xlabel('Number of Tokens')\nplt.ylabel('Frequency')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:48:19.514420Z","iopub.execute_input":"2024-11-25T13:48:19.514779Z","iopub.status.idle":"2024-11-25T13:48:19.757515Z","shell.execute_reply.started":"2024-11-25T13:48:19.514747Z","shell.execute_reply":"2024-11-25T13:48:19.756671Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyzing Query Character Length Distribution","metadata":{}},{"cell_type":"code","source":"# Compute query character lengths\ncharacter_len_list = [len(query) for query in query_list]\nmax_query_length = max(character_len_list)\nprint(\"Maximum query character length:\", max_query_length)\n\n# Plot histogram\nplt.hist(character_len_list)\nplt.title('Query Character Length Distribution')\nplt.xlabel('Number of Characters')\nplt.ylabel('Frequency')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:48:22.782344Z","iopub.execute_input":"2024-11-25T13:48:22.783060Z","iopub.status.idle":"2024-11-25T13:48:23.024461Z","shell.execute_reply.started":"2024-11-25T13:48:22.783025Z","shell.execute_reply":"2024-11-25T13:48:23.023648Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Verifying Sample Counts","metadata":{}},{"cell_type":"code","source":"print(\"Number of samples:\", len(nn_df))\nprint(\"Number of queries:\", len(query_list))\nprint(\"Number of successful queries:\", success_count)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:49:03.181287Z","iopub.execute_input":"2024-11-25T13:49:03.181660Z","iopub.status.idle":"2024-11-25T13:49:03.187502Z","shell.execute_reply.started":"2024-11-25T13:49:03.181627Z","shell.execute_reply":"2024-11-25T13:49:03.186436Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Creating OOF (Out-of-Fold) DataFrame","metadata":{}},{"cell_type":"code","source":"if IS_TRAIN:\n    import pandas as pd\n\n    # Prepare data for OOF DataFrame\n    columns = ['query_list',\n               'true_patent_set_len_list',\n               'score_list',\n               'patent_set_len_list',\n               'neg_count_list',\n               'neg_count_list_2',\n               'query_count_list',\n               'character_len_list']\n    values = np.array([query_list,\n                       true_patent_set_len_list,\n                       score_list,\n                       patent_set_len_list,\n                       neg_count_list,\n                       neg_count_list_2,\n                       query_count_list,\n                       character_len_list]).T\n    \n    oof_df = pd.DataFrame(values, columns=columns)\n    print(oof_df.head())\n    oof_df.to_csv('oof_df.csv', index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:49:05.338569Z","iopub.execute_input":"2024-11-25T13:49:05.338941Z","iopub.status.idle":"2024-11-25T13:49:05.344076Z","shell.execute_reply.started":"2024-11-25T13:49:05.338910Z","shell.execute_reply":"2024-11-25T13:49:05.343052Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Plotting Score vs. True Patent Set Length","metadata":{}},{"cell_type":"code","source":"plt.scatter(true_patent_set_len_list, score_list)\nplt.xlabel('true_patent_set_len_list')\nplt.ylabel('score_list')\nplt.title('Score vs. True Patent Set Length')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:49:11.409092Z","iopub.execute_input":"2024-11-25T13:49:11.409861Z","iopub.status.idle":"2024-11-25T13:49:11.579841Z","shell.execute_reply.started":"2024-11-25T13:49:11.409825Z","shell.execute_reply":"2024-11-25T13:49:11.578860Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Filtering Non-Zero Scores for Analysis","metadata":{}},{"cell_type":"code","source":"# Filter out samples with zero scores\nidxs = [i for i in range(len(score_list)) if score_list[i] != 0]\nfiltered_true_patent_set_len_list = [true_patent_set_len_list[i] for i in idxs]\nfiltered_score_list = [score_list[i] for i in idxs]\nprint(\"Number of samples with non-zero scores:\", len(filtered_true_patent_set_len_list))\n\n# Plotting again with non-zero scores\nplt.scatter(filtered_true_patent_set_len_list, filtered_score_list)\nplt.xlabel('true_patent_set_len_list')\nplt.ylabel('score_list')\nplt.title('Score vs. True Patent Set Length (Non-Zero Scores)')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:49:14.553210Z","iopub.execute_input":"2024-11-25T13:49:14.553559Z","iopub.status.idle":"2024-11-25T13:49:14.704544Z","shell.execute_reply.started":"2024-11-25T13:49:14.553528Z","shell.execute_reply":"2024-11-25T13:49:14.703670Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Handling Empty or Invalid Queries","metadata":{}},{"cell_type":"code","source":"# Define a default query\ndefault_query = 'ti:titonium'\n\n# Replace queries containing 'id:' or empty queries with the default query\nfor i in range(len(query_list)):\n    if 'id:' in query_list[i] or query_list[i].strip() == '':\n        query_list[i] = default_query\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:49:17.823239Z","iopub.execute_input":"2024-11-25T13:49:17.823572Z","iopub.status.idle":"2024-11-25T13:49:17.828399Z","shell.execute_reply.started":"2024-11-25T13:49:17.823545Z","shell.execute_reply":"2024-11-25T13:49:17.827441Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Validating Queries (Optional)","metadata":{}},{"cell_type":"code","source":"if not IS_TRAIN:\n    # Load the test index\n    test_idx = whoosh_utils.load_index('/kaggle/input/uspto-test-index/test_index')\n    searcher = whoosh_utils.get_searcher(test_idx)\n    qp = whoosh_utils.get_query_parser()\n    \n    # Validate queries and replace invalid ones with the default query\n    for i in range(len(query_list)):\n        try:\n            result = whoosh_utils.execute_query(query_list[i], qp, searcher)\n        except Exception as e:\n            print(f\"Error in query {i}: {e}\")\n            query_list[i] = default_query\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:49:21.569206Z","iopub.execute_input":"2024-11-25T13:49:21.570033Z","iopub.status.idle":"2024-11-25T13:50:18.325520Z","shell.execute_reply.started":"2024-11-25T13:49:21.569997Z","shell.execute_reply":"2024-11-25T13:50:18.324561Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Generating Submission File","metadata":{}},{"cell_type":"code","source":"if not IS_TRAIN:\n    import pandas as pd\n\n    # Load the sample submission\n    sub = pd.read_csv('/kaggle/input/uspto-explainable-ai/sample_submission.csv')\n    \n    # Replace the 'query' column with our generated queries\n    sub['query'] = query_list\n    \n    # Save the submission file\n    sub.to_csv('submission.csv', index=False)\n    \n    print(sub.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T13:28:47.432385Z","iopub.status.idle":"2024-11-25T13:28:47.432835Z","shell.execute_reply.started":"2024-11-25T13:28:47.432609Z","shell.execute_reply":"2024-11-25T13:28:47.432632Z"}},"outputs":[],"execution_count":null}]}