{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Dependencies","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.insert(0, \"../input/chenglu-ai4code-source\")","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:38:43.341649Z","iopub.execute_input":"2022-08-11T00:38:43.342927Z","iopub.status.idle":"2022-08-11T00:38:43.378994Z","shell.execute_reply.started":"2022-08-11T00:38:43.342837Z","shell.execute_reply":"2022-08-11T00:38:43.377731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    from nltk.stem import WordNetLemmatizer\nexcept:\n    !pip install ../input/nltk37/nltk/nltk-3.7-py3-none-any.whl\n    !cp -r ../input/nltk37/nltk/nltk_data ~/\n    from nltk.stem import WordNetLemmatizer\nimport pickle\nimport ai4code\nimport os\nimport multiprocessing","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:38:43.382560Z","iopub.execute_input":"2022-08-11T00:38:43.382996Z","iopub.status.idle":"2022-08-11T00:38:55.282158Z","shell.execute_reply.started":"2022-08-11T00:38:43.382956Z","shell.execute_reply":"2022-08-11T00:38:55.280705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"METRIC_SPLIT_MODE = \"max\"\nENSEMBLE_MODE = \"weighted_average\"","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:38:55.283912Z","iopub.execute_input":"2022-08-11T00:38:55.284660Z","iopub.status.idle":"2022-08-11T00:38:55.291281Z","shell.execute_reply.started":"2022-08-11T00:38:55.284609Z","shell.execute_reply":"2022-08-11T00:38:55.289666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    import featurize\n    IN_FEATURIZE = True\n    CODEBERT_PRETRAINED_PATH = \"/home/featurize/codebert-base\"\n    DEBERTA_PRETRAIN_PATH = \"/home/featurize/deberta-v3-small/\"\n    DEBERTA_BASE_PRETRAIN_PATH = \"../input/deberta-v3/deberta-v3-base\"\nexcept:\n    IN_FEATURIZE = False\n    CODEBERT_PRETRAINED_PATH = \"../input/codebert/codebert-base\"\n    DEBERTA_PRETRAIN_PATH = \"../input/deberta-v3/deberta-v3-small\"\n    DEBERTA_BASE_PRETRAIN_PATH = \"../input/deberta-v3/deberta-v3-base\"","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:38:55.292910Z","iopub.execute_input":"2022-08-11T00:38:55.294239Z","iopub.status.idle":"2022-08-11T00:38:55.306670Z","shell.execute_reply.started":"2022-08-11T00:38:55.294194Z","shell.execute_reply":"2022-08-11T00:38:55.305205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"ai4code.cfg = cfg = ai4code.utils.Config(\n    dataset_root = \"../input/AI4Code\" if not IN_FEATURIZE else \"/home/featurize/data/\",\n    encode_files = [\n        {\n            \"name\": \"deberta-v3-v8\",\n            \"preprocessor\": \"preprocessor_v8\",\n            \"pretrain_path\": DEBERTA_PRETRAIN_PATH,\n        },\n#         {\n#             \"name\": \"codebert-v8\",\n#             \"preprocessor\": \"preprocessor_v8\",\n#             \"pretrain_path\": CODEBERT_PRETRAINED_PATH,\n#         }\n    ],\n    checkpoints = [\n#         {\n#             \"path\": \"../input/final-deberta-large-swa/checkpoint_40000_kendall_tau0.9483.pt\",\n#             \"encode\": \"deberta-v3-v8\",\n#             \"pretrain_path\": DEBERTA_BASE_PRETRAIN_PATH,\n#             \"most_common_encodes\": \"../input/debertav3-most-common-encodes/most_common_encodes.pkl\",\n#         },\n\n        {\n            \"path\": \"../input/kafka-deberta-small-2/checkpoint_56102_kendall_tau0.9378.pt\",\n            \"encode\": \"deberta-v3-v8\",\n            \"pretrain_path\": DEBERTA_PRETRAIN_PATH,\n            \"most_common_encodes\": \"../input/debertav3-most-common-encodes/most_common_encodes.pkl\",\n            \"params\": {\n                \"val_anchor_size\": 160\n            }\n        },\n#         {\n#             \"path\": \"../input/kafka-codebert-2/checkpoint_209979_kendall_tau0.9436.pt\",\n#             \"encode\": \"codebert-v8\",\n#             \"pretrain_path\": CODEBERT_PRETRAINED_PATH,\n#             \"most_common_encodes\": None\n#         }\n    ],\n)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:12.943053Z","iopub.execute_input":"2022-08-11T00:43:12.943576Z","iopub.status.idle":"2022-08-11T00:43:12.956041Z","shell.execute_reply.started":"2022-08-11T00:43:12.943542Z","shell.execute_reply":"2022-08-11T00:43:12.954486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ai4code.utils.dump_cfg(cfg)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:14.010521Z","iopub.execute_input":"2022-08-11T00:43:14.010958Z","iopub.status.idle":"2022-08-11T00:43:14.018573Z","shell.execute_reply.started":"2022-08-11T00:43:14.010926Z","shell.execute_reply":"2022-08-11T00:43:14.017159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Tokenization\n\nTokenize all the text in a standalone process to ensure memory been fully released after.","metadata":{}},{"cell_type":"code","source":"%%file /tmp/gen.py\n\nimport pickle\n\ntry:\n    import featurize\n    IN_FEATURIZE = True\n    val_sample_ids = list(pickle.load(open(\"/home/featurize/work/ai4code/data/v8_deberta_small/0.pkl\", \"rb\")).keys())\n    print(\"val sample num: \", len(val_sample_ids))\nexcept Exception as e:\n    IN_FEATURIZE = False\n\nimport sys\nsys.path.insert(0, \"../input/chenglu-ai4code-source\")\nfrom transformers import AutoTokenizer\nimport ai4code\nimport multiprocessing\nfrom tqdm import tqdm\nfrom pathlib import Path\nimport pandas as pd\nfrom functools import reduce\n\nai4code.cfg = ai4code.utils.load_cfg()\ndataset_root = Path(ai4code.cfg.dataset_root)\n\norders = pd.read_csv(dataset_root / \"train_orders.csv\")\nai4code.utils.orders_dict = {}\nfor _, item in orders.iterrows():\n    ai4code.utils.orders_dict[item.id] = item.cell_order.split(\" \")\n\nfor encode_file in ai4code.cfg.encode_files:\n    encode_file[\"tokenizer\"] = AutoTokenizer.from_pretrained(encode_file[\"pretrain_path\"], do_lower_case=True, use_fast=True)\n    encode_file[\"preprocessor\"] = getattr(ai4code.datasets.preprocessor, encode_file[\"preprocessor\"])\n    # dump special tokens at first\n    pickle.dump(dict(\n        hash_id=encode_file[\"tokenizer\"].encode(\"#\", add_special_tokens=False)[0],\n        cls_token_id=encode_file[\"tokenizer\"].cls_token_id,\n        sep_token_id=encode_file[\"tokenizer\"].sep_token_id,\n        pad_token_id=encode_file[\"tokenizer\"].pad_token_id,\n        unk_token_id=encode_file[\"tokenizer\"].unk_token_id,\n    ), open(f\"./special_tokens.{encode_file['name']}.pkl\", \"wb\"))\n\nwith multiprocessing.Pool(processes=multiprocessing.cpu_count()) as pool:\n    if IN_FEATURIZE:\n        json_files = list((dataset_root / \"train\").glob(\"*.json\"))\n        json_files = [json_file for json_file in json_files if json_file.name.split(\".\")[0] in val_sample_ids][:500]\n    else:\n        json_files = list((dataset_root / \"test\").glob(\"*.json\"))\n        # json_files = list((dataset_root / \"train\").glob(\"*.json\"))[:100]\n    print(\"json files total:\", len(json_files))\n\n    results = list(\n        tqdm(\n            pool.imap(ai4code.utils.process, json_files),\n            total=len(json_files),\n        )\n    )\n\nall_data = {sample.id: sample for sample in results}\npickle.dump(all_data, open(\"/tmp/encodes.pkl\", \"wb\"))","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:14.552561Z","iopub.execute_input":"2022-08-11T00:43:14.555145Z","iopub.status.idle":"2022-08-11T00:43:14.565917Z","shell.execute_reply.started":"2022-08-11T00:43:14.555081Z","shell.execute_reply":"2022-08-11T00:43:14.564160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not os.path.exists(\"/tmp/encodes.pkl\"):\n    !python /tmp/gen.py","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:14.757391Z","iopub.execute_input":"2022-08-11T00:43:14.759168Z","iopub.status.idle":"2022-08-11T00:43:14.774171Z","shell.execute_reply.started":"2022-08-11T00:43:14.759081Z","shell.execute_reply":"2022-08-11T00:43:14.770598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"import torch\nimport transformers\nfrom transformers import AutoTokenizer\nimport pandas as pd\nimport multiprocessing\nfrom tqdm import tqdm\nfrom pathlib import Path\nimport os\nimport pickle","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:15.276352Z","iopub.execute_input":"2022-08-11T00:43:15.277504Z","iopub.status.idle":"2022-08-11T00:43:15.285527Z","shell.execute_reply.started":"2022-08-11T00:43:15.277456Z","shell.execute_reply":"2022-08-11T00:43:15.284214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"enable_amp = True if IN_FEATURIZE else False\nbatch_size = 128 if IN_FEATURIZE else 128\nDEVICE = torch.device(\"cuda\") if torch.cuda.is_available() else torch.device(\"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:15.536524Z","iopub.execute_input":"2022-08-11T00:43:15.537578Z","iopub.status.idle":"2022-08-11T00:43:15.545653Z","shell.execute_reply.started":"2022-08-11T00:43:15.537528Z","shell.execute_reply":"2022-08-11T00:43:15.544173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"torch version:\", torch.__version__)\nprint(\"transformers version:\", transformers.__version__)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:16.084455Z","iopub.execute_input":"2022-08-11T00:43:16.085561Z","iopub.status.idle":"2022-08-11T00:43:16.096121Z","shell.execute_reply.started":"2022-08-11T00:43:16.085512Z","shell.execute_reply":"2022-08-11T00:43:16.093889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(\"/tmp/encodes.pkl\", \"rb\") as f:\n    data = pickle.load(f)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:16.547875Z","iopub.execute_input":"2022-08-11T00:43:16.549204Z","iopub.status.idle":"2022-08-11T00:43:16.558944Z","shell.execute_reply.started":"2022-08-11T00:43:16.549169Z","shell.execute_reply":"2022-08-11T00:43:16.557336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TODO: 提前创建 dataset ，创建完后删除 encode 节约内存\ndef create_dataset(checkpoint):\n    state = torch.load(checkpoint['path'])\n    params = {\n        **state[\"params\"],\n        **(checkpoint[\"params\"] if \"params\" in checkpoint else {}),\n    }\n    samples_dir = Path(\"/tmp/samples\") / checkpoint['encode']\n    samples_dir.mkdir(exist_ok=True, parents=True)\n    del state\n    model_encode = checkpoint['encode']\n    special_tokens = ai4code.datasets.SpecialTokenID(\n        **pickle.load(open(f\"./special_tokens.{model_encode}.pkl\", \"rb\"))\n    )\n\n    dataset = ai4code.datasets.MixedDatasetWithSplits(\n        data,\n        special_tokens,\n        only_task_data=True,\n        encode_key=model_encode,\n        anchor_size=params[\"val_anchor_size\"] if \"val_anchor_size\" in params else params[\"anchor_size\"],\n        max_len=params[\"max_len\"],\n        split_len=params[\"split_len\"],\n        global_keywords=checkpoint[\"most_common_encodes\"]\n    )\n    dataset.presist_samples(samples_dir)\n    checkpoint[\"_dataset\"] = dataset\n\nfor checkpoint in ai4code.cfg.checkpoints:\n    create_dataset(checkpoint)\n\nfor checkpoint in ai4code.cfg.checkpoints:\n    # remove all the encode\n    for sample in tqdm(data.values()):\n        sample.cell_encodes[checkpoint['encode']] = None\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:17.062772Z","iopub.execute_input":"2022-08-11T00:43:17.063207Z","iopub.status.idle":"2022-08-11T00:43:17.675440Z","shell.execute_reply.started":"2022-08-11T00:43:17.063175Z","shell.execute_reply":"2022-08-11T00:43:17.673641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_args = (\"with_lstm\", \"with_context_feature\", \"max_len\")\nsubmission_data = []\n\ndef inference(checkpoint, k):\n    model_encode = checkpoint['encode']\n    state = torch.load(checkpoint['path'])\n    params = {\n        **state[\"params\"],\n        **({} if \"params\" not in checkpoint else checkpoint[\"params\"]),\n    }\n    ai4code.utils.print_params(params)\n    model_params = {k: v for k, v in params.items() if k in model_args}\n    model = ai4code.models.MultiHeadModel(checkpoint['pretrain_path'], with_lm=False, dropout=0, **model_params)\n    if \"n_averaged\" in state['model']:\n        state['model'] = {k.replace(\"module.\", \"\"): v for k, v in state['model'].items()}\n    try:\n        model.load_state_dict(state['model'])\n    except Exception as e:\n        print(f\"load model failed with strict mode: {e}, try to load with unstrict\")\n        model.load_state_dict(state['model'], strict=False)\n    model.eval()\n    model.to(DEVICE)\n    del state\n\n    loader = torch.utils.data.DataLoader(\n        checkpoint[\"_dataset\"],\n        num_workers=2,\n        batch_size=batch_size,\n    )\n\n    metric = ai4code.metrics.KendallTauWithSplits(data, params[\"split_len\"], mode=METRIC_SPLIT_MODE)\n\n    with torch.no_grad():\n        for batch in tqdm(loader):\n            ids, mask, targets = [item.to(DEVICE) for item in batch[:3]]\n            sample_ids, cell_keys, split_ids, rank_offset = batch[3:]\n            with torch.cuda.amp.autocast(enabled=enable_amp):\n                in_split, rank, _ = model(ids, mask, lm=False)\n            metric.update((0, in_split, rank, sample_ids, cell_keys, split_ids, rank_offset))\n        print(\"Score: \", metric.compute())\n        with open(f\"/tmp/raw_preds.{k}.pkl\", \"wb\") as f:\n            pickle.dump(metric._raw_preds, f)\n\nfor k, checkpoint in enumerate(ai4code.cfg.checkpoints):\n    inference(checkpoint, k)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:43:18.023557Z","iopub.execute_input":"2022-08-11T00:43:18.023993Z","iopub.status.idle":"2022-08-11T00:43:22.458785Z","shell.execute_reply.started":"2022-08-11T00:43:18.023961Z","shell.execute_reply":"2022-08-11T00:43:22.456447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_raw_preds = [pickle.load(open(f\"/tmp/raw_preds.{k}.pkl\", \"rb\")) for k, _ in enumerate(ai4code.cfg.checkpoints)]\n\nkeys = all_raw_preds[0].keys()\nall_ensemble_preds = {}\n\nfor sample_id, sample in tqdm(data.items()):\n    sample_raw_preds = [raw_preds[sample_id] for raw_preds in all_raw_preds]\n    sample_ensemble_pred = []\n    ensembled_cell_preds = []\n    for cell_preds in zip(*sample_raw_preds):\n        if ENSEMBLE_MODE == \"weighted_average\":\n            cell_pred = list(cell_preds[0])\n            if sample.cell_types[cell_pred[0]] == \"markdown\":\n                scores = torch.sigmoid(torch.tensor([p[2] for p in cell_preds]))\n                ranks = torch.tensor([p[1] for p in cell_preds])\n                weights = scores / scores.sum()\n                rank = (ranks * weights).sum()\n                cell_pred[1] = rank\n            ensembled_cell_preds.append(cell_pred)\n        else:\n            cell_pred = cell_preds[0]\n            if sample.cell_types[cell_pred[0]] == \"markdown\":\n                for cur_cell_pred in cell_preds[1:]:\n                    if cur_cell_pred[2] > cell_pred[2]:\n                        cell_pred = cur_cell_pred\n            ensembled_cell_preds.append(cell_pred)\n    cell_id_predicted = [\n        item[0] for item in sorted(ensembled_cell_preds, key=lambda x: x[1])\n    ]\n    all_ensemble_preds[sample_id] = cell_id_predicted","metadata":{"execution":{"iopub.status.busy":"2022-08-11T00:40:07.099561Z","iopub.execute_input":"2022-08-11T00:40:07.099949Z","iopub.status.idle":"2022-08-11T00:40:07.136750Z","shell.execute_reply.started":"2022-08-11T00:40:07.099914Z","shell.execute_reply":"2022-08-11T00:40:07.135353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if os.path.exists(\"../input/AI4Code/sample_submission.csv\"):\n    sample_df = pd.read_csv(\"../input/AI4Code/sample_submission.csv\")\n    for idx, sample_id in enumerate(sample_df.id):\n        sample_df.loc[idx, 'cell_order'] = \" \".join(all_ensemble_preds[sample_id])\n    sample_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T14:25:32.298084Z","iopub.execute_input":"2022-08-10T14:25:32.298347Z","iopub.status.idle":"2022-08-10T14:25:32.315521Z","shell.execute_reply.started":"2022-08-10T14:25:32.298322Z","shell.execute_reply":"2022-08-10T14:25:32.314617Z"},"trusted":true},"execution_count":null,"outputs":[]}]}