{"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":"# DeBERTa Pre-Training using MLM\n\nI'm fairly new to the domain of NLP or rather machine learning itself. This competition has made me curious about transformers and BERT and is responsible for so many nights of learning these new concepts. This notebook is made from my understanding of Masked Level Modelling (MLM) and probably contains unnecessary/inefficient/incorrect code. Please do let me know if you find any kind of mistakes in this notebook. ","metadata":{}},{"cell_type":"markdown","source":"## Model Training\nModel training has two types: pre-training and fine-tuning.\n\n### Pre-Training\nPre-training is the process of training your model on a large dataset in a **generalized manner** i.e. it can be used for any kind of machine learning problem. \n### Fine-Tuning\nFine-tuning a model simply means to train your model according to a given problem or **particular task**.\n\n### Pre-Training vs Fine-Tuning\nIn pre-training you are actually training the model for any kind of task eg. Classification/Regression, Summarization, Named Entity Recognition (NER), QnA, etc. This is due to the fact that you are training your model to learn about the data in the context of the **domain**. Eg. medical reports, some tweets, or in fact feedback!.\n\nIn fine-tuning you basically take your pretrained model and train it again for some specific task. Now, your fine-tune model has learned about the data in the context of the **domain** as well as that **particular task**. Eg. classifying diseases from medical reports.","metadata":{}},{"cell_type":"markdown","source":"## Pre-Training Techniques\nYou can pre-train your model in two ways:\n1) **MLM** (Masked Language Modelling)\n\n2) **NSP** (Next Sentence Prediction)\n\n### MLM\nMLM is a technique in which you take your tokenized sample and replace some of the tokens with the [MASK] token and train your model with it. The model then tries to predict what should come in the place of that [MASK] token and gradually starts learning about the data.\nMLM teaches the model about the ***relationship between words***\n\nEg. Suppose you have a sentence - 'Deep Learning is so cool! I love neural networks.', now replace few words with the [MASK] token.\n\nMasked Sentence - 'Deep Learning is so [MASK]! I love [MASK] networks.'\n\n### NSP\nIn NSP you input the model with two sentences and your model tries to predict whether the second sentence comes after the first or not. NSP teaches the model about the ***long term dependencies across sentences***.","metadata":{}},{"cell_type":"markdown","source":"This notebook was made after reading manyyy discussions and noticing better performances with MLM. I've not tried it yet personally. Do let me know if this helps you! 😊","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.checkpoint import checkpoint\nfrom transformers import AutoTokenizer, AutoModelWithLMHead\nfrom transformers import AdamW\nfrom tqdm import tqdm\nimport os","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:19:11.277501Z","iopub.execute_input":"2022-08-01T21:19:11.278702Z","iopub.status.idle":"2022-08-01T21:19:13.749292Z","shell.execute_reply.started":"2022-08-01T21:19:11.278589Z","shell.execute_reply":"2022-08-01T21:19:13.748327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    model_name = 'microsoft/deberta-v3-large'\n    epochs = 3\n    batch_size = 4\n    lr = 1e-6\n    weight_decay = 1e-6\n    max_len = 512\n    mask_prob = 0.15  # perc of tokens to convert to mask\n    n_accumulate = 4\n    use_2021 = True\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:19:13.751309Z","iopub.execute_input":"2022-08-01T21:19:13.752211Z","iopub.status.idle":"2022-08-01T21:19:13.793721Z","shell.execute_reply.started":"2022-08-01T21:19:13.752173Z","shell.execute_reply":"2022-08-01T21:19:13.791979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=CFG.seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed()","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:19:13.795150Z","iopub.execute_input":"2022-08-01T21:19:13.796242Z","iopub.status.idle":"2022-08-01T21:19:13.804920Z","shell.execute_reply.started":"2022-08-01T21:19:13.796202Z","shell.execute_reply":"2022-08-01T21:19:13.804027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.use_2021:\n    competition_path = \"../input/feedback-prize-2021/\"\n    df = pd.read_csv('../input/feedback-pseudo-labelling-full-2021-dataset/train_2021_preds.csv');\n    df = df[df['in_2022'] == False]\nelse:\n    competition_path = \"../input/feedback-prize-effectiveness/\"\n    df = pd.read_csv(competition_path + 'train.csv')","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:19:13.808111Z","iopub.execute_input":"2022-08-01T21:19:13.808574Z","iopub.status.idle":"2022-08-01T21:19:14.537324Z","shell.execute_reply.started":"2022-08-01T21:19:13.808509Z","shell.execute_reply":"2022-08-01T21:19:14.536384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fetch_essay_texts(df, train=True):\n    if train:\n        base_path = competition_path + 'train/'\n    else:\n        base_path = competition_path + 'test/'\n        \n    essay_texts = {}\n    for filename in os.listdir(base_path):\n        with open(base_path + filename) as f:\n            text = f.readlines()\n            full_text = ' '.join([x for x in text])\n            essay_text = ' '.join([x for x in full_text.split()])\n        essay_texts[filename[:-4]] = essay_text\n    df['essay_text'] = [essay_texts[essay_id] for essay_id in df['essay_id'].values]\n    return df\n\nfetch_essay_texts(df)","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:19:14.539347Z","iopub.execute_input":"2022-08-01T21:19:14.539739Z","iopub.status.idle":"2022-08-01T21:19:22.711955Z","shell.execute_reply.started":"2022-08-01T21:19:14.539703Z","shell.execute_reply":"2022-08-01T21:19:22.710731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenizer = AutoTokenizer.from_pretrained(CFG.model_name)\nmodel = AutoModelWithLMHead.from_pretrained(CFG.model_name)","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:19:22.716801Z","iopub.execute_input":"2022-08-01T21:19:22.719487Z","iopub.status.idle":"2022-08-01T21:19:39.108665Z","shell.execute_reply.started":"2022-08-01T21:19:22.719450Z","shell.execute_reply":"2022-08-01T21:19:39.107775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"special_tokens = tokenizer.encode_plus('[CLS] [SEP] [MASK] [PAD]',\n                                        add_special_tokens = False,\n                                        return_tensors='pt')\nspecial_tokens = torch.flatten(special_tokens[\"input_ids\"])\nprint(special_tokens)","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:19:39.110029Z","iopub.execute_input":"2022-08-01T21:19:39.110493Z","iopub.status.idle":"2022-08-01T21:19:39.118243Z","shell.execute_reply.started":"2022-08-01T21:19:39.110455Z","shell.execute_reply":"2022-08-01T21:19:39.116827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def getMaskedLabels(input_ids):\n    rand = torch.rand(input_ids.shape)\n    mask_arr = (rand < CFG.mask_prob)\n    # Preventing special tokens to get replace by the [MASK] token\n    for special_token in special_tokens:\n        token = special_token.item()\n        mask_arr *= (input_ids != token)\n    selection = torch.flatten(mask_arr[0].nonzero()).tolist()\n    input_ids[selection] = 128000\n    \n    return input_ids","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:19:39.119639Z","iopub.execute_input":"2022-08-01T21:19:39.120218Z","iopub.status.idle":"2022-08-01T21:19:39.130465Z","shell.execute_reply.started":"2022-08-01T21:19:39.120180Z","shell.execute_reply":"2022-08-01T21:19:39.129515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MLMDataset:\n    def __init__(self, data, tokenizer):\n        self.data = data\n        self.tokenizer = tokenizer\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        text = self.data[idx]\n        \n        tokenized_data = self.tokenizer.encode_plus(\n                            text,\n                            max_length = CFG.max_len,\n                            truncation = True,\n                            padding = 'max_length',\n                            add_special_tokens = True,\n                            return_tensors = 'pt'\n                        )\n        input_ids = torch.flatten(tokenized_data.input_ids)\n        attention_mask = torch.flatten(tokenized_data.attention_mask)\n        labels = getMaskedLabels(input_ids)\n        \n        return {\n            'input_ids': input_ids,\n            'attention_mask': attention_mask,\n            'labels': labels\n        }","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:19:39.132138Z","iopub.execute_input":"2022-08-01T21:19:39.132576Z","iopub.status.idle":"2022-08-01T21:19:39.142147Z","shell.execute_reply.started":"2022-08-01T21:19:39.132484Z","shell.execute_reply":"2022-08-01T21:19:39.140867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"essay_data = df[\"essay_text\"].unique()\ndataset = MLMDataset(essay_data, tokenizer)\ndataloader = DataLoader(dataset, batch_size=CFG.batch_size, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:21:10.944950Z","iopub.execute_input":"2022-08-01T21:21:10.945302Z","iopub.status.idle":"2022-08-01T21:21:11.298356Z","shell.execute_reply.started":"2022-08-01T21:21:10.945274Z","shell.execute_reply":"2022-08-01T21:21:11.297392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(df), len(essay_data)","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:21:11.300261Z","iopub.execute_input":"2022-08-01T21:21:11.300727Z","iopub.status.idle":"2022-08-01T21:21:11.312823Z","shell.execute_reply.started":"2022-08-01T21:21:11.300687Z","shell.execute_reply":"2022-08-01T21:21:11.311018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:21:11.314533Z","iopub.execute_input":"2022-08-01T21:21:11.314918Z","iopub.status.idle":"2022-08-01T21:21:11.340356Z","shell.execute_reply.started":"2022-08-01T21:21:11.314883Z","shell.execute_reply":"2022-08-01T21:21:11.338795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_loop(model, device):\n    model.train()\n    batch_losses = []\n    loop = tqdm(dataloader, leave=True)\n    for batch_num, batch in enumerate(loop):\n        optimizer.zero_grad()\n        input_ids = batch[\"input_ids\"].to(device)\n        attention_mask = batch[\"attention_mask\"].to(device)\n        labels = batch[\"labels\"].to(device)\n\n        outputs = model(input_ids, attention_mask=attention_mask, labels=labels)\n\n        loss = outputs.loss\n        batch_loss = loss / CFG.n_accumulate\n        batch_losses.append(batch_loss.item())\n    \n        loop.set_description(f\"Epoch {epoch + 1}\")\n        loop.set_postfix(loss=batch_loss.item())\n        batch_loss.backward()\n        \n        if batch_num % CFG.n_accumulate == 0 or batch_num == len(dataloader):\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5)\n            optimizer.step()\n            model.zero_grad()\n\n    return np.mean(batch_losses)","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:21:11.401598Z","iopub.execute_input":"2022-08-01T21:21:11.402056Z","iopub.status.idle":"2022-08-01T21:21:11.416989Z","shell.execute_reply.started":"2022-08-01T21:21:11.402029Z","shell.execute_reply":"2022-08-01T21:21:11.415972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\nfrom kaggle_secrets import UserSecretsClient\n\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"WANDB\")\nwandb.login(key=secret_value_0)\nwandb.init(project='feedback-prize-effectiveness', name='mlm-aurora-large')","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:21:11.650279Z","iopub.execute_input":"2022-08-01T21:21:11.650937Z","iopub.status.idle":"2022-08-01T21:21:21.436291Z","shell.execute_reply.started":"2022-08-01T21:21:11.650903Z","shell.execute_reply":"2022-08-01T21:21:21.435361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = CFG.device\nmodel.to(device)\nhistory = []\nbest_loss = np.inf\nprev_loss = np.inf\nmodel.gradient_checkpointing_enable()\nprint(f\"Gradient Checkpointing: {model.is_gradient_checkpointing}\")\n\nfor epoch in range(CFG.epochs):\n    loss = train_loop(model, device)\n    history.append(loss)\n    print(f\"Loss: {loss}\")\n    if loss < best_loss:\n        print(\"New Best Loss {:.4f} -> {:.4f}, Saving Model\".format(prev_loss, loss))\n        # torch.save(model.state_dict(), \"./deberta_mlm.pt\")\n        model.save_pretrained('./')\n        best_loss = loss\n    prev_loss = loss","metadata":{"execution":{"iopub.status.busy":"2022-08-01T21:21:21.439304Z","iopub.execute_input":"2022-08-01T21:21:21.439623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}