{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":59575,"databundleVersionId":8060720,"sourceType":"competition"},{"sourceId":8479599,"sourceType":"datasetVersion","datasetId":4517815},{"sourceId":8553271,"sourceType":"datasetVersion","datasetId":5109610},{"sourceId":174185912,"sourceType":"kernelVersion"}],"dockerImageVersionId":30747,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# T5 based doc2query (titles)\n※ I'm sorry if my English is poor, I'm just practicing. <br>\n\nIn the field of information retrieval, there is a task to estimate queries that are likely to be retrieved from documents (doc2query).<br>\ndoc2query is used for document extensions and contributes to the improvement of recall in information retrieval.\n\nIn this case, I'll apply the query estimated by doc2query to USPTO task instead of document extension.\n\n### Reference Link\n- Doc2query: Document Expansion by Query Prediction: https://arxiv.org/abs/1904.08375","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/whoosh-wheel-2-7-4/Whoosh-2.7.4-py2.py3-none-any.whl","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import polars as pl\nimport pickle\nfrom pathlib import Path\nimport re\nimport numpy as np\nfrom numpy.typing import NDArray\nimport polars as pl\nfrom tqdm import tqdm\nfrom typing import Any\nimport whoosh_utils\n\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nfrom transformers import T5Tokenizer, T5ForConditionalGeneration, Trainer, TrainingArguments\nimport datasets\nfrom  collections import Counter","metadata":{"execution":{"iopub.status.busy":"2024-07-17T16:26:48.32024Z","iopub.execute_input":"2024-07-17T16:26:48.321718Z","iopub.status.idle":"2024-07-17T16:26:48.600427Z","shell.execute_reply.started":"2024-07-17T16:26:48.321678Z","shell.execute_reply":"2024-07-17T16:26:48.599311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"comp_data_dir = Path(\"/kaggle/input/uspto-explainable-ai\")\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare Data\nTo learn model for doc2query, a pair of document and retrieved query is needed.\n\nThe following notebooks collected pairs of queries and titles instead of documents.<br>\n\n\nBased on https://www.kaggle.com/code/tubotubo/uspto-simulated-annealing-baseline","metadata":{}},{"cell_type":"markdown","source":"### Create Index for collecting Data","metadata":{}},{"cell_type":"code","source":"%%python\nfrom pathlib import Path\n\nimport polars as pl\nfrom tqdm import tqdm\n\nimport whoosh_utils\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\ntest_nn = pl.read_csv(comp_data_dir / \"nearest_neighbors.csv\").head(2500).lazy()\n\n\n# Filtering only the patent meta-information that appears in the test\nall_pub = test_nn.melt().collect().get_column(\"value\").unique()\nneg_random = meta.filter(pl.col(\"publication_number\").is_in(all_pub).not_()).sample(75000)\nmeta = meta.filter(pl.col(\"publication_number\").is_in(all_pub))\n\nmeta = pl.concat([meta, neg_random])\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([\"claims\", \"description\"]))\n    patents.append(patent)\npatent: pl.LazyFrame = pl.concat(patents)\npatent = patent.with_columns(\n    pl.lit(\"\").alias(\"claims\"),\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(\"dev_meta_with_text.parquet\")\n\n# create index\ndocuments = meta_with_text.to_dicts()\n\ntest_nn.collect().write_csv(\"dev.csv\")\nPath(\"dev_index\").mkdir(parents=True, exist_ok=True)\nwhoosh_utils.create_index(\"dev_index\", documents)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dev = pl.read_csv(\"dev.csv\")\nmeta = pl.read_parquet(\"dev_meta_with_text.parquet\")\n\n# dev index\ndev_idx = whoosh_utils.load_index(\"./dev_index\")\nsearcher = whoosh_utils.get_searcher(dev_idx)\nqp = whoosh_utils.get_query_parser()\n\n\nwith open(\"/kaggle/input/uspto-ti-cpc-tfidf/tfidf.pkl\", \"rb\") as f:\n    ti_tfidf = pickle.load(f)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Get Query and Titiles¶\n- Query Conditions\n    - Only numbers are excluded from the query\n    - Words with a frequency of less than 5 are also excluded from the query\n- Titles Conditions\n    - Target topK titles tied to the query (K=5)","metadata":{}},{"cell_type":"code","source":"NUMBER_REGEX = re.compile(r'^(\\d+|\\d{1,3}(,\\d{3})*)(\\.\\d+)?$')\n\ntmp = meta.select(\"title\").with_columns(pl.col(\"title\").str.split(\" \").alias(\"word\")).explode(\"word\").group_by(\"word\").count()\nword2freq_lut = { w: c for w,c in tmp.to_numpy()}\n\nvocab = ti_tfidf.get_feature_names_out()\n\nqueries = [word for word in vocab \n         if not(NUMBER_REGEX.match(word)) \n         and word in word2freq_lut \n         and word2freq_lut[word] >= 5]\nlen(queries)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"k = 5\n\nquery_title_pair = pl.DataFrame({\"query\": [], \"title\": []}).cast({\"query\": pl.String, \"title\": pl.String})\n\nfor query in tqdm(queries):\n    ti_query = f\"ti:{query}\"\n    cand = whoosh_utils.execute_query(ti_query, qp, searcher)\n    topk_cand = cand[:k]\n    \n    meta_topk_cand = (\n        meta.filter(pl.col(\"publication_number\").is_in(topk_cand))\n        .with_columns(pl.Series(np.array([query]*len(topk_cand))).alias(\"query\"))\n    )\n    if len(meta_topk_cand) != 0:\n        query_title_pair = pl.concat([query_title_pair, meta_topk_cand.select(\"query\", \"title\")])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"query_title_pair","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"query_title_pair.write_csv(\"/kaggle/working/query_titles.csv\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model\nI use T5, which is commonly used in doc2query.\n![](https://cdn-ak.f.st-hatena.com/images/fotolife/l/lib-arts/20191125/20191125191712.png)\n\n### Reference Link\n- https://github.com/castorini/docTTTTTquery\n- Exploring the Limits of Transfer Learning with a Unified Text-to-Text Transformer: https://arxiv.org/pdf/1910.10683","metadata":{}},{"cell_type":"code","source":"tokenizer = T5Tokenizer.from_pretrained('castorini/doc2query-t5-base-msmarco')\nmodel = T5ForConditionalGeneration.from_pretrained('castorini/doc2query-t5-base-msmarco')\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-07-17T16:26:48.760878Z","iopub.execute_input":"2024-07-17T16:26:48.761337Z","iopub.status.idle":"2024-07-17T16:26:57.5493Z","shell.execute_reply.started":"2024-07-17T16:26:48.7613Z","shell.execute_reply":"2024-07-17T16:26:57.54802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# データの前処理関数\ndef preprocess_function(examples):\n    inputs = [doc for doc in examples[\"title\"]]\n    model_inputs = tokenizer(inputs, max_length=128, truncation=True, padding=\"max_length\")\n    # ラベルのトークナイズ\n    with tokenizer.as_target_tokenizer():\n        labels = tokenizer([str(label) for label in examples[\"query\"]], max_length=128, truncation=True, padding=\"max_length\")\n\n    model_inputs[\"labels\"] = labels[\"input_ids\"]\n    return model_inputs","metadata":{"execution":{"iopub.status.busy":"2024-07-17T16:56:35.125096Z","iopub.execute_input":"2024-07-17T16:56:35.12596Z","iopub.status.idle":"2024-07-17T16:56:35.133231Z","shell.execute_reply.started":"2024-07-17T16:56:35.125924Z","shell.execute_reply":"2024-07-17T16:56:35.13197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train, val = train_test_split(query_title_pair, test_size=0.2, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2024-07-17T16:26:57.562668Z","iopub.execute_input":"2024-07-17T16:26:57.563182Z","iopub.status.idle":"2024-07-17T16:26:57.667704Z","shell.execute_reply.started":"2024-07-17T16:26:57.563147Z","shell.execute_reply":"2024-07-17T16:26:57.666558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train =  datasets.Dataset.from_pandas(train.to_pandas())\nds_val =  datasets.Dataset.from_pandas(val.to_pandas())","metadata":{"execution":{"iopub.status.busy":"2024-07-17T16:38:34.348137Z","iopub.execute_input":"2024-07-17T16:38:34.348619Z","iopub.status.idle":"2024-07-17T16:38:34.501837Z","shell.execute_reply.started":"2024-07-17T16:38:34.348585Z","shell.execute_reply":"2024-07-17T16:38:34.500664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenized_ds_train = ds_train.map(preprocess_function, batched=True)\ntokenized_ds_val = ds_val.map(preprocess_function, batched=True)","metadata":{"execution":{"iopub.status.busy":"2024-07-17T16:56:39.734841Z","iopub.execute_input":"2024-07-17T16:56:39.735278Z","iopub.status.idle":"2024-07-17T16:56:55.212772Z","shell.execute_reply.started":"2024-07-17T16:56:39.735245Z","shell.execute_reply":"2024-07-17T16:56:55.211598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenized_ds_val","metadata":{"execution":{"iopub.status.busy":"2024-07-17T16:56:56.536149Z","iopub.execute_input":"2024-07-17T16:56:56.536714Z","iopub.status.idle":"2024-07-17T16:56:56.545593Z","shell.execute_reply.started":"2024-07-17T16:56:56.536671Z","shell.execute_reply":"2024-07-17T16:56:56.544142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# トレーニングの設定\ntraining_args = TrainingArguments(\n    output_dir=\"/kaggle/working/model\",\n    evaluation_strategy=\"epoch\",\n    learning_rate=5e-5,\n    per_device_train_batch_size=16,\n    per_device_eval_batch_size=16,\n    num_train_epochs=3,\n    weight_decay=0.01,\n    push_to_hub=False,\n    report_to=[],\n    save_total_limit=3,\n)","metadata":{"execution":{"iopub.status.busy":"2024-07-17T16:56:58.214306Z","iopub.execute_input":"2024-07-17T16:56:58.214737Z","iopub.status.idle":"2024-07-17T16:56:58.222401Z","shell.execute_reply.started":"2024-07-17T16:56:58.214703Z","shell.execute_reply":"2024-07-17T16:56:58.221295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    model=model,\n    args=training_args,\n    train_dataset=tokenized_ds_train,\n    eval_dataset=tokenized_ds_val,\n    tokenizer=tokenizer,\n)\ntrainer.train()","metadata":{"execution":{"iopub.status.busy":"2024-07-17T16:57:00.199816Z","iopub.execute_input":"2024-07-17T16:57:00.200585Z","iopub.status.idle":"2024-07-17T16:58:03.423508Z","shell.execute_reply.started":"2024-07-17T16:57:00.20055Z","shell.execute_reply":"2024-07-17T16:58:03.421504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference\nIf you are submitting a submission, save the model, turn off the Internet, and change the model path accordingly!\n※ Score submission Error","metadata":{}},{"cell_type":"code","source":"# model_path = \"/kaggle/input/uspto-t5-based-doc2query/model/XXXX\"\n# tokenizer = T5Tokenizer.from_pretrained(model_path)\n# model = T5ForConditionalGeneration.from_pretrained(model_path)\n# model.to(device)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%python\nfrom pathlib import Path\n\nimport polars as pl\nfrom tqdm import tqdm\n\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\ntest_nn = pl.scan_csv(comp_data_dir / \"test.csv\")\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\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    if year is None or month is None:\n        patent_path = comp_data_dir / f\"patent_data/nan_nan.parquet\"\n    else:\n        patent_path = comp_data_dir / f\"patent_data/{year}_{month}.parquet\"\n    patent = pl.scan_parquet(patent_path).select(pl.exclude([\"claims\", \"description\"]))\n#     patent = pl.scan_parquet(patent_path).select(pl.exclude([\"claims\"])).with_columns(pl.col(\"description\").str.slice(0,1000))\n    patents.append(patent)\npatent: pl.LazyFrame = pl.concat(patents)\npatent = patent.with_columns(\n    pl.lit(\"\").alias(\"claims\"),\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\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = pl.read_csv(comp_data_dir / \"test.csv\")\ntest_meta = pl.read_parquet(\"meta_with_text.parquet\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def generate_text(input_text, max_length=1024, N=3):\n    input_ids = tokenizer.encode(input_text, return_tensors=\"pt\", truncation=True).to(device)\n    output_ids = model.generate(input_ids, \n                                max_length=max_length, \n                                num_return_sequences=N,\n                                num_beams=5, \n                                early_stopping=True)\n    \n    output_texts = [tokenizer.decode(ids, skip_special_tokens=True) for ids in output_ids]\n    return output_texts\n\ndef generate_query(titles, k=10):\n    outputs = []\n    for title in titles:\n        output = generate_text(title)\n        outputs.extend(output)\n    word_cnt = Counter(outputs) # 頻度が高い順\n    \n    words = []\n    try:\n        for word, i in word_cnt.most_common():\n            if len(word.split()) != 1:\n                continue\n            words.append(f\"ti:{word}\")\n            if len(words) > k:\n                break\n        return \" OR \".join(words)\n    except:\n        return \"ti:device\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results = []\n\n\nfor i in tqdm(range(len(test))):\n    target = test[i].to_numpy().flatten()[1:].tolist()\n    meta_i = test_meta.filter(pl.col(\"publication_number\").is_in(target))\n\n    if len(meta_i) == 0:\n        # append dummy\n        results.append({\"publication_number\": test[i, \"publication_number\"], \"query\": \"ti:device\"})\n        print(\"\\t Append Dummy\", i)\n        continue\n    titles = meta_i.get_column(\"title\").fill_null(\"\")\n    query = generate_query(titles)\n    results.append({\"publication_number\": test[i, \"publication_number\"], \"query\": query})","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf /kaggle/working/*","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pl.DataFrame(results)\nsubmission.write_csv(\"submission.csv\")\n\nsubmission","metadata":{},"execution_count":null,"outputs":[]}]}