{"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":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, \n# but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-04T08:57:34.833493Z","iopub.execute_input":"2022-08-04T08:57:34.834240Z","iopub.status.idle":"2022-08-04T08:57:34.847921Z","shell.execute_reply.started":"2022-08-04T08:57:34.834148Z","shell.execute_reply":"2022-08-04T08:57:34.846528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ndevice=torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:57:41.588776Z","iopub.execute_input":"2022-08-04T08:57:41.589166Z","iopub.status.idle":"2022-08-04T08:57:42.217952Z","shell.execute_reply.started":"2022-08-04T08:57:41.589133Z","shell.execute_reply":"2022-08-04T08:57:42.216891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df=pd.read_csv(\"../input/contradictory-my-dear-watson/train.csv\")\ntest_df=pd.read_csv(\"../input/contradictory-my-dear-watson/test.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:57:43.875373Z","iopub.execute_input":"2022-08-04T08:57:43.876526Z","iopub.status.idle":"2022-08-04T08:57:43.964812Z","shell.execute_reply.started":"2022-08-04T08:57:43.876487Z","shell.execute_reply":"2022-08-04T08:57:43.963777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df=train_df.sample(frac=1.0)\ntrain_df=train_df.iloc[2000:, :]\nval_df=train_df.iloc[:2000, :]","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:57:44.187704Z","iopub.execute_input":"2022-08-04T08:57:44.188214Z","iopub.status.idle":"2022-08-04T08:57:44.205159Z","shell.execute_reply.started":"2022-08-04T08:57:44.188173Z","shell.execute_reply":"2022-08-04T08:57:44.204083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from datasets import Dataset\ntrain_ds=Dataset.from_pandas(train_df)\nval_ds=Dataset.from_pandas(val_df)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:57:45.038373Z","iopub.execute_input":"2022-08-04T08:57:45.039362Z","iopub.status.idle":"2022-08-04T08:57:45.437909Z","shell.execute_reply.started":"2022-08-04T08:57:45.039324Z","shell.execute_reply":"2022-08-04T08:57:45.436572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoTokenizer\ntokenizer = AutoTokenizer.from_pretrained('xlm-roberta-base')","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:57:45.843481Z","iopub.execute_input":"2022-08-04T08:57:45.844968Z","iopub.status.idle":"2022-08-04T08:57:48.227568Z","shell.execute_reply.started":"2022-08-04T08:57:45.844923Z","shell.execute_reply":"2022-08-04T08:57:48.226559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_cols = ['label']\n\nfor part in ['premise', 'hypothesis']:\n    train_ds = train_ds.map(\n        lambda x: tokenizer(\n            x[part], max_length=128, padding='max_length',\n            truncation=True\n        ), batched=True\n    )\n    for col in ['input_ids', 'attention_mask']:\n        train_ds = train_ds.rename_column(\n            col, part+'_'+col\n        )\n        all_cols.append(part+'_'+col)\nprint(all_cols)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:57:48.229465Z","iopub.execute_input":"2022-08-04T08:57:48.229832Z","iopub.status.idle":"2022-08-04T08:57:52.413757Z","shell.execute_reply.started":"2022-08-04T08:57:48.229794Z","shell.execute_reply":"2022-08-04T08:57:52.412727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_cols = ['label']\n\nfor part in ['premise', 'hypothesis']:\n    val_ds = val_ds.map(\n        lambda x: tokenizer(\n            x[part], max_length=128, padding='max_length',\n            truncation=True\n        ), batched=True\n    )\n    for col in ['input_ids', 'attention_mask']:\n        val_ds = val_ds.rename_column(\n            col, part+'_'+col\n        )\n        all_cols.append(part+'_'+col)\nprint(all_cols)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:57:52.415172Z","iopub.execute_input":"2022-08-04T08:57:52.415626Z","iopub.status.idle":"2022-08-04T08:57:53.263738Z","shell.execute_reply.started":"2022-08-04T08:57:52.415586Z","shell.execute_reply":"2022-08-04T08:57:53.262740Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ntrain_ds.set_format(type='torch', columns=all_cols)\nval_ds.set_format(type=\"torch\", columns=all_cols)\n# initialize the dataloader\nbatch_size = 16\ntrain_loader = torch.utils.data.DataLoader(\n    train_ds, batch_size=batch_size, shuffle=True\n)\nval_loader = torch.utils.data.DataLoader(\n    val_ds, batch_size=batch_size, shuffle=False\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:57:53.266265Z","iopub.execute_input":"2022-08-04T08:57:53.266663Z","iopub.status.idle":"2022-08-04T08:57:53.274384Z","shell.execute_reply.started":"2022-08-04T08:57:53.266626Z","shell.execute_reply":"2022-08-04T08:57:53.273307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoModel\nbert=AutoModel.from_pretrained(\"xlm-roberta-base\").to(device)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:57:53.275754Z","iopub.execute_input":"2022-08-04T08:57:53.276282Z","iopub.status.idle":"2022-08-04T08:58:00.863315Z","shell.execute_reply.started":"2022-08-04T08:57:53.276245Z","shell.execute_reply":"2022-08-04T08:58:00.862057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nclass MyModel(nn.Module):\n    def __init__(self, bert, embedding_dim=768, num_classes=3):\n        super(MyModel, self).__init__()\n        self.bert=bert\n        self.linear=nn.Linear(embedding_dim*3, num_classes)\n    \n    # define mean pooling function\n    def mean_pool(self, token_embeds, attention_mask):\n        # reshape attention_mask to cover 768-dimension embeddings\n        in_mask = attention_mask.unsqueeze(-1).expand(\n            token_embeds.size()\n        ).float()\n        # perform mean-pooling but exclude padding tokens (specified by in_mask)\n        pool = torch.sum(token_embeds * in_mask, 1) / torch.clamp(\n            in_mask.sum(1), min=1e-9\n        )\n        return pool\n    \n    def forward(self, batch_data):\n        u=self.bert(batch_data['premise_input_ids'].to(device), batch_data['premise_attention_mask'].to(device), output_hidden_states=True)\n        u=self.mean_pool(u.last_hidden_state, batch_data['premise_attention_mask'].to(device))\n        v=self.bert(batch_data['hypothesis_input_ids'].to(device), batch_data['hypothesis_attention_mask'].to(device), output_hidden_states=True)\n        v=self.mean_pool(v.last_hidden_state, batch_data['hypothesis_attention_mask'].to(device))\n        x=torch.cat([u, v, torch.abs(u-v)], dim=-1)\n        y=self.linear(x)\n        return y","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:58:03.846541Z","iopub.execute_input":"2022-08-04T08:58:03.846965Z","iopub.status.idle":"2022-08-04T08:58:03.914680Z","shell.execute_reply.started":"2022-08-04T08:58:03.846927Z","shell.execute_reply":"2022-08-04T08:58:03.913247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_fn=torch.nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2022-08-04T08:58:22.643724Z","iopub.execute_input":"2022-08-04T08:58:22.644181Z","iopub.status.idle":"2022-08-04T08:58:22.649769Z","shell.execute_reply.started":"2022-08-04T08:58:22.644145Z","shell.execute_reply":"2022-08-04T08:58:22.648291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model=MyModel(bert).to(device)\noptim = torch.optim.Adam(model.parameters(), lr=2e-5)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T09:25:36.616403Z","iopub.execute_input":"2022-08-04T09:25:36.616886Z","iopub.status.idle":"2022-08-04T09:25:36.648149Z","shell.execute_reply.started":"2022-08-04T09:25:36.616843Z","shell.execute_reply":"2022-08-04T09:25:36.647084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers.optimization import get_linear_schedule_with_warmup\ntotal_steps = int(len(train_ds) / batch_size)\nwarmup_steps = int(0.1 * total_steps)\nscheduler = get_linear_schedule_with_warmup(optim, num_warmup_steps=warmup_steps,num_training_steps=total_steps - warmup_steps)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T09:25:36.994510Z","iopub.execute_input":"2022-08-04T09:25:36.995028Z","iopub.status.idle":"2022-08-04T09:25:37.007580Z","shell.execute_reply.started":"2022-08-04T09:25:36.994980Z","shell.execute_reply":"2022-08-04T09:25:37.006279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs=1\nfrom tqdm import tqdm\nfrom sklearn.metrics import accuracy_score, f1_score\nfor i in range(num_epochs):\n    for batch in tqdm(train_loader):\n        optim.zero_grad()\n        y=model(batch)\n        loss=loss_fn(y, batch['label'].to(device))\n        loss.backward()\n        optim.step()\n        scheduler.step()\n    correct=0\n    total=0\n    epoch_loss=0\n    with torch.no_grad():\n        for batch in tqdm(train_loader):\n            y=model(batch)\n            loss=loss_fn(y, batch['label'].to(device))\n            epoch_loss+=loss.item()\n            preds=y.argmax(-1)\n            correct+=(preds==batch['label'].to(device)).sum()\n            total+=preds.shape[0]\n        print(f\"Epoch: {i} Train accuracy: {correct/total*100} Loss: {epoch_loss/total}\")\n    \n    correct=0\n    total=0\n    epoch_loss=0\n    with torch.no_grad():\n        for batch in tqdm(val_loader):\n            y=model(batch)\n            loss=loss_fn(y, batch['label'].to(device))\n            epoch_loss+=loss.item()\n            preds=y.argmax(-1)\n            correct+=(preds==batch['label'].to(device)).sum()\n            total+=preds.shape[0]\n        print(f\"Epoch: {i} Val accuracy: {correct/total*100} Loss: {epoch_loss/total}\")","metadata":{"execution":{"iopub.status.busy":"2022-08-04T09:25:37.448455Z","iopub.execute_input":"2022-08-04T09:25:37.449527Z","iopub.status.idle":"2022-08-04T09:31:38.683654Z","shell.execute_reply.started":"2022-08-04T09:25:37.449490Z","shell.execute_reply":"2022-08-04T09:31:38.682489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_cols = []\ntest_ds=Dataset.from_pandas(test_df)\nfor part in ['premise', 'hypothesis']:\n    test_ds = test_ds.map(\n        lambda x: tokenizer(\n            x[part], max_length=128, padding='max_length',\n            truncation=True\n        ), batched=True\n    )\n    for col in ['input_ids', 'attention_mask']:\n        test_ds = test_ds.rename_column(\n            col, part+'_'+col\n        )\n        all_cols.append(part+'_'+col)\nprint(all_cols)\n\ntest_ds.set_format(type=\"torch\", columns=all_cols)\n# initialize the dataloader\nbatch_size = 16\ntest_loader = torch.utils.data.DataLoader(\n    test_ds, batch_size=batch_size, shuffle=False\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T09:32:27.697974Z","iopub.execute_input":"2022-08-04T09:32:27.698810Z","iopub.status.idle":"2022-08-04T09:32:29.895714Z","shell.execute_reply.started":"2022-08-04T09:32:27.698770Z","shell.execute_reply":"2022-08-04T09:32:29.894720Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ids=test_df['id']\noutputs=[]\nwith torch.no_grad():\n    for batch in tqdm(test_loader):\n        y=model(batch)\n        preds=y.argmax(-1)\n        outputs.extend(list(preds.detach().cpu().numpy()))","metadata":{"execution":{"iopub.status.busy":"2022-08-04T09:32:34.859533Z","iopub.execute_input":"2022-08-04T09:32:34.859920Z","iopub.status.idle":"2022-08-04T09:33:16.864349Z","shell.execute_reply.started":"2022-08-04T09:32:34.859886Z","shell.execute_reply":"2022-08-04T09:33:16.863122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(outputs)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T09:33:30.655544Z","iopub.execute_input":"2022-08-04T09:33:30.655952Z","iopub.status.idle":"2022-08-04T09:33:30.662463Z","shell.execute_reply.started":"2022-08-04T09:33:30.655919Z","shell.execute_reply":"2022-08-04T09:33:30.661461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df=pd.DataFrame({\"id\":ids, \"prediction\":outputs})\nsub_df.set_index(\"id\", inplace=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-04T09:33:47.032259Z","iopub.execute_input":"2022-08-04T09:33:47.032990Z","iopub.status.idle":"2022-08-04T09:33:47.045117Z","shell.execute_reply.started":"2022-08-04T09:33:47.032952Z","shell.execute_reply":"2022-08-04T09:33:47.044107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df","metadata":{"execution":{"iopub.status.busy":"2022-08-04T09:33:48.857010Z","iopub.execute_input":"2022-08-04T09:33:48.857791Z","iopub.status.idle":"2022-08-04T09:33:48.869864Z","shell.execute_reply.started":"2022-08-04T09:33:48.857752Z","shell.execute_reply":"2022-08-04T09:33:48.868510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv(\"submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-08-04T09:33:51.865957Z","iopub.execute_input":"2022-08-04T09:33:51.867014Z","iopub.status.idle":"2022-08-04T09:33:51.881225Z","shell.execute_reply.started":"2022-08-04T09:33:51.866966Z","shell.execute_reply":"2022-08-04T09:33:51.880312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}