{"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":"tpu1vmV38","dataSources":[{"sourceId":59575,"databundleVersionId":8060720,"sourceType":"competition"},{"sourceId":8479599,"sourceType":"datasetVersion","datasetId":4517815},{"sourceId":8723242,"sourceType":"datasetVersion","datasetId":5234797},{"sourceId":174185912,"sourceType":"kernelVersion"},{"sourceId":187077654,"sourceType":"kernelVersion"}],"dockerImageVersionId":30664,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport glob\nfrom tqdm import tqdm\nimport gc\nfrom collections import defaultdict\nimport datetime\nimport json\nimport itertools\nimport matplotlib.pyplot as plt\nimport copy\nimport math\nfrom collections import Counter\nimport numpy as np\nimport gc\nimport time\nimport random\n\nimport whoosh_utils","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install polars","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 実験内容  \nOOMしないように、TPU sessionでindexを作る\n\n","metadata":{}},{"cell_type":"code","source":"IS_TRAIN = True\nIS_REDUCE_TRAIN = True\n\nIS_ADD_NEG_0_PATTERN = True","metadata":{"execution":{"iopub.status.busy":"2024-07-06T09:14:41.829792Z","iopub.execute_input":"2024-07-06T09:14:41.830408Z","iopub.status.idle":"2024-07-06T09:14:41.835930Z","shell.execute_reply.started":"2024-07-06T09:14:41.830368Z","shell.execute_reply":"2024-07-06T09:14:41.834824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nPUB_MAX_COUNT = 1000\nSUB_MAX_COUNT = 300\nR_MAX = 2\nUPDATE_MIN = 1\nPATTERN_NUM_MAX = 10000\nbeam_width = 50\nIS_BEAM_SORTED = True\n\"\"\"\n\nNEG_MAX_COUNT = 100\nPUB_MAX_COUNT = 50 + NEG_MAX_COUNT # 10\nSUB_MAX_COUNT = 50 + NEG_MAX_COUNT # 100\n\nR_MAX = 2\nUPDATE_MIN = 1\nPATTERN_NUM_MAX = 1000\n\nIS_BEAM_SORTED = True # いらない beamsearchのときは必要","metadata":{"execution":{"iopub.status.busy":"2024-07-06T09:14:41.837203Z","iopub.execute_input":"2024-07-06T09:14:41.837600Z","iopub.status.idle":"2024-07-06T09:14:41.849962Z","shell.execute_reply.started":"2024-07-06T09:14:41.837573Z","shell.execute_reply":"2024-07-06T09:14:41.848922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"T0 = SUB_MAX_COUNT # 300\nT1 = SUB_MAX_COUNT / 10 # 30\n\nmax_time = 5 # 5","metadata":{"execution":{"iopub.status.busy":"2024-07-06T09:14:41.851929Z","iopub.execute_input":"2024-07-06T09:14:41.852352Z","iopub.status.idle":"2024-07-06T09:14:41.862028Z","shell.execute_reply.started":"2024-07-06T09:14:41.852322Z","shell.execute_reply":"2024-07-06T09:14:41.861016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nimport pickle\n# base_path = '/kaggle/input/upsto-cpc-baseline-all-feature-low-memory/'\n# base_path = '/kaggle/input/precompute-cpc-title-desc-max-10000/'\n# base_path = '/kaggle/input/precompute-cpc-title-desc-max-10000-word-to-int/'\nbase_path = '/kaggle/input/precompute-claim-set-max-1000-add-claims/'\n\n\"\"\"\nword_to_number = pickle.load(open(base_path + 'word_to_number.pkl', 'rb'))\nnumber_to_word = pickle.load(open(base_path + 'number_to_word.pkl', 'rb'))\nword_to_pub_set = pickle.load(open(base_path + 'cpc_to_pub_set.pkl', 'rb'))\nword_to_pubcount = pickle.load(open(base_path + 'cpc_to_pubcount.pkl', 'rb'))\n\"\"\"\npublication_to_word = pickle.load(open(base_path + 'publication_to_word.pkl', 'rb'))","metadata":{"execution":{"iopub.status.busy":"2024-07-06T09:14:41.863374Z","iopub.execute_input":"2024-07-06T09:14:41.863765Z","iopub.status.idle":"2024-07-06T09:15:22.001904Z","shell.execute_reply.started":"2024-07-06T09:14:41.863729Z","shell.execute_reply":"2024-07-06T09:15:22.000697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IS_TRAIN:\n\n    nn_df = pd.read_csv('/kaggle/input/uspto-explainable-ai/nearest_neighbors.csv', nrows=4000)\n\n    # target50個が全て1975年以上のものを使用する\n    use_idxs = []\n    for i, pub in enumerate(nn_df['publication_number'].values):\n        if pub not in publication_to_word:\n            continue\n        \n        flag = True\n        for pub2 in nn_df.values[i, 1:]:\n            # if pub2 not in publication_to_publication_date:\n            if pub2 not in publication_to_word:\n                flag = False\n                break\n        if flag:\n            use_idxs.append(i)\n\n    print(len(use_idxs))\n    nn_df = nn_df.loc[use_idxs, :].reset_index(drop=True)\n    nn_df = nn_df.head(2500)\n    \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')\n    pass\n","metadata":{"execution":{"iopub.status.busy":"2024-07-06T09:15:22.003363Z","iopub.execute_input":"2024-07-06T09:15:22.003696Z","iopub.status.idle":"2024-07-06T09:15:22.411376Z","shell.execute_reply.started":"2024-07-06T09:15:22.003667Z","shell.execute_reply":"2024-07-06T09:15:22.410216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%python\nfrom pathlib import Path\n\nimport polars as pl\nfrom tqdm import tqdm\n\nimport whoosh_utils\n\nIS_TRAIN = True\n\n# 乱数固定\nimport os\n# import torch\nimport numpy as np\nimport random\ndef set_seed(seed=None, cudnn_deterministic=True):\n    if seed is None:\n        seed = 42\n\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n    # torch.manual_seed(seed)\n    # torch.cuda.manual_seed(seed)\n    # torch.backends.cudnn.deterministic = cudnn_deterministic  # A100,effnetだとFalseの方が早い\n    # torch.backends.cudnn.benchmark = False\nset_seed()\n    \ncomp_data_dir = Path(\"/kaggle/input/uspto-explainable-ai\")\n\n# Read patent since 1975\nmeta = pl.scan_parquet(comp_data_dir / \"patent_metadata.parquet\")\nmeta = (\n    meta.with_columns(\n        pl.col(\"publication_date\").dt.year().alias(\"year\"),\n        pl.col(\"publication_date\").dt.month().alias(\"month\"),\n    )\n    .filter(pl.col(\"publication_date\") >= pl.date(1975, 1, 1))\n    .rename({\"cpc_codes\": \"cpc\"})\n    .collect()\n)\n\n\n\nif IS_TRAIN:\n    # test_nn = pl.scan_csv(comp_data_dir / \"test.csv\")\n    # test_nn = pl.read_csv(comp_data_dir / \"nearest_neighbors.csv\", n_rows=4000).head(2500).lazy()\n    # test_nn = pl.read_csv(comp_data_dir / \"nearest_neighbors.csv\", n_rows=4000).head(2500)\n    test_nn = pl.read_csv('nn_df_for_index.csv')\n    \n    test_nn_pub = test_nn.melt().get_column(\"value\").unique()\n    # 候補数を合わせるためにnegative sampleを抽出 候補数: (2500 + 1500) * 50 = 200,000\n    neg_sample = (\n        pl.read_csv(comp_data_dir / \"nearest_neighbors.csv\")\n        .filter(~pl.col(\"publication_number\").is_in(test_nn_pub))\n        .sample(1500, seed=42)\n    )\n    test_nn = pl.concat([test_nn, neg_sample]).lazy()\nelse:\n    test_nn = pl.read_csv(comp_data_dir / \"test.csv\")\n    # test_nn = pl.read_csv(comp_data_dir / \"nearest_neighbors.csv\", n_rows=4000).head(2500).lazy()\n    \n    test_nn_pub = test_nn.melt().get_column(\"value\").unique()\n    # 候補数を合わせるためにnegative sampleを抽出 候補数: (2500 + 1500) * 50 = 200,000\n    neg_sample = (\n        pl.read_csv(comp_data_dir / \"nearest_neighbors.csv\")\n        .filter(~pl.col(\"publication_number\").is_in(test_nn_pub))\n        .sample(1500, seed=42)\n        .rename({f'neighbor_{i}':f'target_{i}' for i in range(50)})\n    )\n    test_nn = pl.concat([test_nn, neg_sample]).lazy()\n\n# Filtering only the patent meta-information that appears in the test\nall_pub = test_nn.melt().collect().get_column(\"value\").unique()\nmeta = meta.filter(pl.col(\"publication_number\").is_in(all_pub))\n\nprint('all_pub_len', len(all_pub))\n\n# Join meta information\npatents = []\nn_unique = meta.select([\"year\", \"month\"]).n_unique()\nfor (year, month), _ in tqdm(meta.group_by([\"year\", \"month\"]), total=n_unique):\n    patent_path = comp_data_dir / f\"patent_data/{year}_{month}.parquet\"\n    patent = pl.scan_parquet(patent_path).select(pl.exclude([\"description\"]))\n    patents.append(patent)\npatent: pl.LazyFrame = pl.concat(patents)\npatent = patent.with_columns(\n    pl.lit(\"\").alias(\"description\"),\n)\nmeta_with_text = (\n    meta.lazy().join(patent, on=\"publication_number\", how=\"left\").collect(streaming=True)\n)\nmeta_with_text.write_parquet(\"meta_with_text.parquet\")\n\n# create index\ndocuments = meta_with_text.to_dicts()\nPath(\"test_index\").mkdir(parents=True, exist_ok=True)\nwhoosh_utils.create_index(\"test_index\", documents)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T09:15:22.412940Z","iopub.execute_input":"2024-07-06T09:15:22.413354Z","iopub.status.idle":"2024-07-06T09:16:56.459485Z","shell.execute_reply.started":"2024-07-06T09:15:22.413315Z","shell.execute_reply":"2024-07-06T09:16:56.457544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}