{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":10737,"databundleVersionId":290346,"sourceType":"competition"},{"sourceId":1658223,"sourceType":"datasetVersion","datasetId":823501}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Introduction: BERT using PyTorch\nBERT (Bidirectional Encoder Representations from Transformers) is a state-of-the-art NLP model developed by Google that understands context in a sentence by leveraging bidirectional attention. It has revolutionized NLP tasks like text classification, question answering, and named entity recognition.\n\nIn this notebook, we will explore how to implement BERT using PyTorch, fine-tune it on a dataset, and evaluate its performance. We will leverage Hugging Face's Transformers library to efficiently train and deploy the model. Let's dive into BERT’s power and applications!","metadata":{}},{"cell_type":"markdown","source":"## Importing Libraries and Setting the Configurations","metadata":{}},{"cell_type":"code","source":"# Preprocessing base libraries\nimport numpy as np\nimport pandas as pd\n\n# For visualization\nimport seaborn as sns\nfrom matplotlib import rc\nfrom pylab import rcParams\nimport matplotlib.pyplot as plt\n\n# For preprocessing the text data\nimport re\nfrom textwrap import wrap\nfrom textblob import TextBlob\nfrom collections import Counter\nfrom string import punctuation\nfrom nltk import word_tokenize, ngrams\n\n# Model evaluation and processing libraries\nfrom sklearn.utils import resample\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix, classification_report, f1_score\n\n# PyTorch required libraries\nimport torch\nfrom torch import nn, optim\nfrom torch.utils.data import Dataset, DataLoader\n\n# Transformer libraries \nimport transformers\nfrom transformers import BertModel, BertTokenizer, AdamW, get_linear_schedule_with_warmup\n\nimport time\nimport warnings\nfrom tqdm import tqdm\nfrom collections import defaultdict\nwarnings.filterwarnings('ignore')\n\n# Setting the Plotting Defaults\n%matplotlib inline\n%config InlineBackend.figure_format = 'retina'\nsns.set(style = 'whitegrid', palette = 'muted', font_scale = 1.2)\nHAPPY_COLORS_PALETTE = [\"#01BEFE\", \"#FFDD00\", \"#FF7D00\", \"#FF006D\", \"#ADFF02\", \"#8F00FF\"]\nsns.set_palette(sns.color_palette(HAPPY_COLORS_PALETTE))\nrcParams['figure.figsize'] = 12, 8\n\n# seeds for reproducibility\nRANDOM_SEED = 42\nnp.random.seed(RANDOM_SEED)\ntorch.manual_seed(RANDOM_SEED)\n\n# Using GPU if available otherwise CPU Cores \ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint (device)\n\n# Punctuations\npunctuation = punctuation + '\"\"“”’' + '∞θ÷α•à−β∅³π‘₹´°£€\\×™√²—–&'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:38:40.562005Z","iopub.execute_input":"2025-02-04T20:38:40.562301Z","iopub.status.idle":"2025-02-04T20:38:40.578252Z","shell.execute_reply.started":"2025-02-04T20:38:40.562279Z","shell.execute_reply":"2025-02-04T20:38:40.577544Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Importing the dataset\ndf = pd.read_csv('/kaggle/input/quora-insincere-questions-classification/train.csv')\nprint (\"The Shape of the dataset : \", df.shape)\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:38:42.430544Z","iopub.execute_input":"2025-02-04T20:38:42.430889Z","iopub.status.idle":"2025-02-04T20:38:46.457277Z","shell.execute_reply.started":"2025-02-04T20:38:42.430859Z","shell.execute_reply":"2025-02-04T20:38:46.456389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df['target'].value_counts() # The Target is highly skewed as we can see","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:38:46.458228Z","iopub.execute_input":"2025-02-04T20:38:46.458456Z","iopub.status.idle":"2025-02-04T20:38:46.474735Z","shell.execute_reply.started":"2025-02-04T20:38:46.458437Z","shell.execute_reply":"2025-02-04T20:38:46.473911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Checing for the null values\ndf.isnull().sum(axis = 0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:38:46.476556Z","iopub.execute_input":"2025-02-04T20:38:46.476838Z","iopub.status.idle":"2025-02-04T20:38:46.618520Z","shell.execute_reply.started":"2025-02-04T20:38:46.476818Z","shell.execute_reply":"2025-02-04T20:38:46.617860Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Exploratory Data Analysis","metadata":{}},{"cell_type":"code","source":"# Changing the target values for the EDA for now\ndf['target'] = df['target'].map({1 : 'Insincere', 0 : 'Sincere'})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:38:46.619555Z","iopub.execute_input":"2025-02-04T20:38:46.619872Z","iopub.status.idle":"2025-02-04T20:38:46.647174Z","shell.execute_reply.started":"2025-02-04T20:38:46.619845Z","shell.execute_reply":"2025-02-04T20:38:46.646560Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Distribution of Insincere and Sincere Questions","metadata":{}},{"cell_type":"code","source":"# Distribution of the Insincere and Sincere Questions\nplt.rcParams['figure.figsize'] = [10, 5]\n\n# Creating the labels for the piechart\ntypes = df['target'].value_counts()\nlabels = list(types.index)\naggregate = list(types.values)\npercentage = [(x*100)/sum(aggregate) for x in aggregate]\nprint (\"The percentages of Sincere and Insincere Questions are : \", percentage)\n\n# Plotting the Piechart to see the percentage distribution of the questions\nplt.rcParams.update({'font.size': 16})\nexplode = (0, 0.1)\nplt.pie(aggregate, labels = labels, autopct='%1.2f%%', shadow=True, colors = ['darkcyan', 'crimson'])\nplt.legend(labels, loc = 'best')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:38:46.647985Z","iopub.execute_input":"2025-02-04T20:38:46.648285Z","iopub.status.idle":"2025-02-04T20:38:47.002450Z","shell.execute_reply.started":"2025-02-04T20:38:46.648254Z","shell.execute_reply":"2025-02-04T20:38:47.001634Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Ngrams (Unigrams and Bigrams)","metadata":{}},{"cell_type":"code","source":"# Copying the dataset for the ngrams so that it won't effect the further processes\ndf_ = df.copy()\ndf_['question_text'] = df_['question_text'].str.replace(r'[^\\w\\d\\s]',' ')\n\n# Segregating the questions\ndf_insincere = \" \".join(df_.loc[df_.target == 'Insincere', 'question_text'])\ndf_sincere = \" \".join(df_.loc[df_.target == 'Sincere', 'question_text'])\n\n# Tokenizing the Sentences\ntokenized_insincere = word_tokenize(df_insincere)\ntokenized_sincere = word_tokenize(df_sincere)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:38:47.003487Z","iopub.execute_input":"2025-02-04T20:38:47.003822Z","iopub.status.idle":"2025-02-04T20:40:54.870975Z","shell.execute_reply.started":"2025-02-04T20:38:47.003789Z","shell.execute_reply":"2025-02-04T20:40:54.870251Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Unigrams\nunigram_insincere = ngrams(tokenized_insincere, 1)\nunigram_sincere = ngrams(tokenized_sincere, 1)\n\n# Making the Frequency chart for the Unigrams\nfrequency_insincere = Counter(unigram_insincere) \nfrequency_sincere = Counter(unigram_sincere)\n\ndf_freq_insincere = pd.DataFrame(frequency_insincere.most_common(20))\ndf_freq_sincere = pd.DataFrame(frequency_sincere.most_common(20))\n\n# Barplot that shows the top 20 Unigrams\nplt.rcParams['figure.figsize'] = [20, 12]\nfig, ax = plt.subplots(1, 2)\nsns.set(font_scale = 1.3, style = 'darkgrid')\n\nsns_sincere = sns.barplot(x = df_freq_sincere[1], y = df_freq_sincere[0], color = 'darkslateblue', ax = ax[0])\nsns_insincere = sns.barplot(x = df_freq_insincere[1], y = df_freq_insincere[0], color = 'cadetblue', ax = ax[1])\n\n# Setting axes\nsns_sincere.set(title = \"Top 20 Unigrams of the Sincere Questions\", ylabel = \"Unigrams\", xlabel = \"Frequency\")\nsns_insincere.set(title = \"Top 20 Unigrams of the Insincere Questions\", xlabel = \"Frequency\", ylabel = \"\");","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:40:54.871845Z","iopub.execute_input":"2025-02-04T20:40:54.872066Z","iopub.status.idle":"2025-02-04T20:41:03.219448Z","shell.execute_reply.started":"2025-02-04T20:40:54.872046Z","shell.execute_reply":"2025-02-04T20:41:03.218601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Bigrams\nbigram_insincere = ngrams(tokenized_insincere, 2)\nbigram_sincere = ngrams(tokenized_sincere, 2)\n\n# Making the Frequency chart for the Bigrams\nfrequency_insincere = Counter(bigram_insincere) \nfrequency_sincere = Counter(bigram_sincere)\n\ndf_freq_insincere = pd.DataFrame(frequency_insincere.most_common(20))\ndf_freq_sincere = pd.DataFrame(frequency_sincere.most_common(20))\n\n# Barplot that shows the top 20 Bigrams\nplt.rcParams['figure.figsize'] = [20, 12]\nfig, ax = plt.subplots(1, 2)\nsns.set(font_scale = 1.3, style = 'darkgrid')\n\nsns_sincere = sns.barplot(x = df_freq_sincere[1], y = df_freq_sincere[0], color = 'darkslateblue', ax = ax[0])\nsns_insincere = sns.barplot(x = df_freq_insincere[1], y = df_freq_insincere[0], color = 'cadetblue', ax = ax[1])\n\n# Setting axes\nsns_sincere.set(title = \"Top 20 Bigrams of the Sincere Questions\", ylabel = \"Bigrams\", xlabel = \"Frequency\")\nsns_insincere.set(title = \"Top 20 Bigrams of the Insincere Questions\", xlabel = \"Frequency\", ylabel = \"\");","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:41:03.222321Z","iopub.execute_input":"2025-02-04T20:41:03.222677Z","iopub.status.idle":"2025-02-04T20:41:16.441464Z","shell.execute_reply.started":"2025-02-04T20:41:03.222643Z","shell.execute_reply":"2025-02-04T20:41:16.440599Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Distribution of Polarity and Subjectivity of the questions","metadata":{}},{"cell_type":"code","source":"# Calculating Polarity of the questions \ndef sentiment_polarity(questions):\n    # Sentiment polarity of the questions\n    pol = []\n    for i in questions:\n        analysis = TextBlob(i)\n        pol.append(analysis.sentiment.polarity)\n    return pol\n\n# Subjectivity of the questions\ndef sentiment_subjectivity(questions):\n    # Sentiment subjectivity of the questions\n    sub = []\n    for i in questions:\n        analysis = TextBlob(i)\n        sub.append(analysis.sentiment.subjectivity)\n    return sub\n\n# Appeding the polarity and subjectivity of the text in the dataframe\ndf['polarity'] = sentiment_polarity(df['question_text'])\ndf['subjectivity'] = sentiment_subjectivity(df['question_text'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:41:16.442720Z","iopub.execute_input":"2025-02-04T20:41:16.443040Z","iopub.status.idle":"2025-02-04T20:48:28.154041Z","shell.execute_reply.started":"2025-02-04T20:41:16.443012Z","shell.execute_reply":"2025-02-04T20:48:28.153273Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Distribution plot\nplt.rcParams['figure.figsize'] = [20, 10]\n# sns.set(style = 'white', font_scale = 1.5)\n\ndist_ = sns.distplot(df['polarity'], kde = False, color = 'deepskyblue')\n# Setting the axes\ndist_.set(title = 'Distribution of the Polarity', xlabel = 'Polarity', ylabel = 'Frequency');","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:48:28.154963Z","iopub.execute_input":"2025-02-04T20:48:28.155282Z","iopub.status.idle":"2025-02-04T20:48:28.843571Z","shell.execute_reply.started":"2025-02-04T20:48:28.155252Z","shell.execute_reply":"2025-02-04T20:48:28.842703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Distribution plot\nplt.rcParams['figure.figsize'] = [20, 10]\ndist_ = sns.distplot(df['subjectivity'], kde = False, color = 'red')\n# Setting the axes\ndist_.set(title = 'Distribution of the Subjectivity', xlabel = 'Subjectivity', ylabel = 'Frequency');","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:48:28.844614Z","iopub.execute_input":"2025-02-04T20:48:28.844974Z","iopub.status.idle":"2025-02-04T20:48:29.518330Z","shell.execute_reply.started":"2025-02-04T20:48:28.844941Z","shell.execute_reply":"2025-02-04T20:48:29.517491Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Preprocessing","metadata":{}},{"cell_type":"code","source":"# Changing the target values for BERT\ndf['target'] = df['target'].map({'Insincere' : 1, 'Sincere' : 0})\n\n# Dropping polarity and subjectivity as it's not going to be used in training BERT\ndf.drop(columns = ['polarity', 'subjectivity'], inplace = True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:48:29.519346Z","iopub.execute_input":"2025-02-04T20:48:29.519701Z","iopub.status.idle":"2025-02-04T20:48:29.598979Z","shell.execute_reply.started":"2025-02-04T20:48:29.519665Z","shell.execute_reply":"2025-02-04T20:48:29.598300Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Contraction Dictionary for the expansion\ncontractions_dict = {\n    \"ain't\": \"am not\", \"aren't\": \"are not\", \"can't\": \"cannot\", \"can't've\": \"cannot have\", \"'cause\": \"because\",\n    \"could've\": \"could have\", \"couldn't\": \"could not\", \"couldn't've\": \"could not have\", \"didn't\": \"did not\", \"doesn't\": \"does not\",\n    \"doesn’t\": \"does not\", \"don't\": \"do not\", \"don’t\": \"do not\", \"hadn't\": \"had not\", \"hadn't've\": \"had not have\", \"hasn't\": \"has not\",\n    \"haven't\": \"have not\", \"he'd\": \"he had\", \"he'd've\": \"he would have\", \"he'll\": \"he will\", \"he'll've\": \"he will have\", \"he's\": \"he is\",\n    \"how'd\": \"how did\", \"how'd'y\": \"how do you\", \"how'll\": \"how will\", \"how's\": \"how is\", \"i'd\": \"i would\", \"i'd've\": \"i would have\",\n    \"i'll\": \"i will\", \"i'll've\": \"i will have\", \"i'm\": \"i am\", \"i've\": \"i have\", \"isn't\": \"is not\", \"it'd\": \"it would\", \"it'd've\": \"it would have\",\n    \"it'll\": \"it will\", \"it'll've\": \"it will have\", \"it's\": \"it is\", \"let's\": \"let us\", \"ma'am\": \"madam\", \"mayn't\": \"may not\",\"might've\": \"might have\",\n    \"mightn't\": \"might not\", \"mightn't've\": \"might not have\", \"must've\": \"must have\", \"mustn't\": \"must not\", \"mustn't've\": \"must not have\",\n    \"needn't\": \"need not\", \"needn't've\": \"need not have\", \"o'clock\": \"of the clock\", \"oughtn't\": \"ought not\", \"oughtn't've\": \"ought not have\",\n    \"shan't\": \"shall not\",\"sha'n't\": \"shall not\", \"shan't've\": \"shall not have\", \"she'd\": \"she would\", \"she'd've\": \"she would have\",\n    \"she'll\": \"she will\", \"she'll've\": \"she will have\", \"she's\": \"she is\", \"should've\": \"should have\", \"shouldn't\": \"should not\",\n    \"shouldn't've\": \"should not have\", \"so've\": \"so have\", \"so's\": \"so is\", \"that'd\": \"that would\", \"that'd've\": \"that would have\",\n    \"that's\": \"that is\", \"there'd\": \"there would\", \"there'd've\": \"there would have\", \"there's\": \"there is\", \"they'd\": \"they would\",\n    \"they'd've\": \"they would have\", \"they'll\": \"they will\", \"they'll've\": \"they will have\", \"they're\": \"they are\", \"they've\": \"they have\",\n    \"to've\": \"to have\", \"wasn't\": \"was not\", \"we'd\": \"we would\", \"we'd've\": \"we would have\", \"we'll\": \"we will\", \"we'll've\": \"we will have\",\n    \"we're\": \"we are\", \"we've\": \"we have\", \"weren't\": \"were not\", \"what'll\": \"what will\", \"what'll've\": \"what will have\", \"what're\": \"what are\",\n    \"what's\": \"what is\", \"what've\": \"what have\", \"when's\": \"when is\", \"when've\": \"when have\", \"where'd\": \"where did\", \"where's\": \"where is\",\n    \"where've\": \"where have\", \"who'll\": \"who will\", \"who'll've\": \"who will have\", \"who's\": \"who is\", \"who've\": \"who have\", \"why's\": \"why is\",\n    \"why've\": \"why have\", \"will've\": \"will have\", \"won't\": \"will not\", \"won't've\": \"will not have\", \"would've\": \"would have\",\n    \"wouldn't\": \"would not\", \"wouldn't've\": \"would not have\", \"y'all\": \"you all\", \"y’all\": \"you all\", \"y'all'd\": \"you all would\",\n    \"y'all'd've\": \"you all would have\", \"y'all're\": \"you all are\", \"y'all've\": \"you all have\", \"you'd\": \"you would\", \"you'd've\": \"you would have\",\n    \"you'll\": \"you will\", \"you'll've\": \"you will have\", \"you're\": \"you are\", \"you've\": \"you have\", \"ain’t\": \"am not\", \"aren’t\": \"are not\",\n    \"can’t\": \"cannot\", \"can’t’ve\": \"cannot have\", \"’cause\": \"because\", \"could’ve\": \"could have\", \"couldn’t\": \"could not\", \"couldn’t’ve\": \"could not have\",\n    \"didn’t\": \"did not\", \"doesn’t\": \"does not\", \"don’t\": \"do not\", \"don’t\": \"do not\", \"hadn’t\": \"had not\", \"hadn’t’ve\": \"had not have\",\n    \"hasn’t\": \"has not\", \"haven’t\": \"have not\", \"he’d\": \"he had\", \"he’d’ve\": \"he would have\", \"he’ll\": \"he will\", \"he’ll’ve\": \"he will have\",\n    \"he’s\": \"he is\", \"how’d\": \"how did\", \"how’d’y\": \"how do you\", \"how’ll\": \"how will\", \"how’s\": \"how is\", \"i’d\": \"i would\", \"i’d’ve\": \"i would have\",\n    \"i’ll\": \"i will\", \"i’ll’ve\": \"i will have\", \"i’m\": \"i am\", \"i’ve\": \"i have\", \"isn’t\": \"is not\", \"it’d\": \"it would\", \"it’d’ve\": \"it would have\",\n    \"it’ll\": \"it will\", \"it’ll’ve\": \"it will have\", \"it’s\": \"it is\", \"let’s\": \"let us\", \"ma’am\": \"madam\", \"mayn’t\": \"may not\",\n    \"might’ve\": \"might have\", \"mightn’t\": \"might not\", \"mightn’t’ve\": \"might not have\", \"must’ve\": \"must have\", \"mustn’t\": \"must not\",\n    \"mustn’t’ve\": \"must not have\", \"needn’t\": \"need not\", \"needn’t’ve\": \"need not have\", \"o’clock\": \"of the clock\",\n    \"oughtn’t\": \"ought not\", \"oughtn’t’ve\": \"ought not have\", \"shan’t\": \"shall not\", \"sha’n’t\": \"shall not\", \"shan’t’ve\": \"shall not have\",\n    \"she’d\": \"she would\", \"she’d’ve\": \"she would have\", \"she’ll\": \"she will\", \"she’ll’ve\": \"she will have\", \"she’s\": \"she is\",\n    \"should’ve\": \"should have\", \"shouldn’t\": \"should not\", \"shouldn’t’ve\": \"should not have\", \"so’ve\": \"so have\", \"so’s\": \"so is\",\n    \"that’d\": \"that would\", \"that’d’ve\": \"that would have\", \"that’s\": \"that is\", \"there’d\": \"there would\", \"there’d’ve\": \"there would have\",\n    \"there’s\": \"there is\", \"they’d\": \"they would\", \"they’d’ve\": \"they would have\", \"they’ll\": \"they will\", \"they’ll’ve\": \"they will have\",\n    \"they’re\": \"they are\", \"they’ve\": \"they have\", \"to’ve\": \"to have\", \"wasn’t\": \"was not\", \"we’d\": \"we would\", \"we’d’ve\": \"we would have\",\n    \"we’ll\": \"we will\", \"we’ll’ve\": \"we will have\", \"we’re\": \"we are\", \"we’ve\": \"we have\", \"weren’t\": \"were not\", \"what’ll\": \"what will\",\n    \"what’ll’ve\": \"what will have\", \"what’re\": \"what are\", \"what’s\": \"what is\", \"what’ve\": \"what have\", \"when’s\": \"when is\",\n    \"when’ve\": \"when have\", \"where’d\": \"where did\", \"where’s\": \"where is\", \"where’ve\": \"where have\", \"who’ll\": \"who will\",\n    \"who’ll’ve\": \"who will have\", \"who’s\": \"who is\", \"who’ve\": \"who have\",\"why’s\": \"why is\", \"why’ve\": \"why have\", \"will’ve\": \"will have\",\n    \"won’t\": \"will not\", \"won’t’ve\": \"will not have\", \"would’ve\": \"would have\", \"wouldn’t\": \"would not\", \"wouldn’t’ve\": \"would not have\",\n    \"y’all\": \"you all\", \"y’all\": \"you all\", \"y’all’d\": \"you all would\", \"y’all’d’ve\": \"you all would have\", \"y’all’re\": \"you all are\",\n    \"y’all’ve\": \"you all have\", \"you’d\": \"you would\", \"you’d’ve\": \"you would have\", \"you’ll\": \"you will\", \"you’ll’ve\": \"you will have\",\n    \"you’re\": \"you are\", \"you’re\": \"you are\", \"you’ve\": \"you have\"\n}\n\ncontractions_re = re.compile('(%s)' % '|'.join(contractions_dict.keys()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:48:29.599836Z","iopub.execute_input":"2025-02-04T20:48:29.600139Z","iopub.status.idle":"2025-02-04T20:48:29.617415Z","shell.execute_reply.started":"2025-02-04T20:48:29.600117Z","shell.execute_reply":"2025-02-04T20:48:29.616601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to clean the html from the Questions\ndef cleanhtml(raw_html):\n    cleanr = re.compile('<.*?>')\n    cleantext = re.sub(cleanr, '', raw_html)\n    return cleantext\n\n# Function expand the contractions if there's any\ndef expand_contractions(s, contractions_dict = contractions_dict):\n    def replace(match):\n        return contractions_dict[match.group(0)]\n    return contractions_re.sub(replace, s)\n\n# Function to preprocess the questions\ndef main_preprocessing(question):\n    global question_sent\n    \n    # Removing the HTML\n    question = question.apply(lambda x: cleanhtml(x))\n    \n    # Removing the email ids\n    question = question.apply(lambda x: re.sub('\\S+@\\S+','', x))\n    \n    # Removing The URLS\n    question = question.apply(lambda x: re.sub(\"((http\\://|https\\://|ftp\\://)|(www.))+(([a-zA-Z0-9\\.-]+\\.[a-zA-Z]{2,4})|([0-9]{1,3}\\.[0-9]{1,3}\\.[0-9]{1,3}\\.[0-9]{1,3}))(/[a-zA-Z0-9%:/-_\\?\\.'~]*)?\",'', x))\n    \n    # Mapping the contractions\n    question = question.apply(lambda x: expand_contractions(x))\n    \n    # Stripping the possessives\n    question = question.apply(lambda x: x.replace(\"'s\", ''))\n    question = question.apply(lambda x: x.replace('’s', ''))\n    question = question.apply(lambda x: x.replace(\"\\'s\", ''))\n    question = question.apply(lambda x: x.replace(\"\\’s\", ''))\n    \n    # Removing the Trailing and leading whitespace and double spaces\n    question = question.apply(lambda x: re.sub(' +', ' ',x))\n    \n    # Removing punctuations from the question\n    question = question.apply(lambda x: ''.join(word for word in x if word not in punctuation))\n    \n    # Removing the Trailing and leading whitespace and double spaces again as removing punctuation might lead to a white space\n    question = question.apply(lambda x: re.sub(' +', ' ',x))\n    \n    return question","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:48:29.618320Z","iopub.execute_input":"2025-02-04T20:48:29.618601Z","iopub.status.idle":"2025-02-04T20:48:29.632993Z","shell.execute_reply.started":"2025-02-04T20:48:29.618581Z","shell.execute_reply":"2025-02-04T20:48:29.632256Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Doing some preprocessing\ndf['processed_text'] = main_preprocessing(df['question_text'])\ndf['processed_text'] = df['processed_text'].astype('str')\n\ndf = resample(df, random_state = RANDOM_SEED) # Resampling","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:48:29.633690Z","iopub.execute_input":"2025-02-04T20:48:29.633926Z","iopub.status.idle":"2025-02-04T20:49:42.218944Z","shell.execute_reply.started":"2025-02-04T20:48:29.633894Z","shell.execute_reply":"2025-02-04T20:49:42.218227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.head(20)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T20:49:42.219776Z","iopub.execute_input":"2025-02-04T20:49:42.220083Z","iopub.status.idle":"2025-02-04T20:49:42.229772Z","shell.execute_reply.started":"2025-02-04T20:49:42.220052Z","shell.execute_reply":"2025-02-04T20:49:42.228862Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## BERT Modelling","metadata":{}},{"cell_type":"markdown","source":"#### BERT Configurations","metadata":{}},{"cell_type":"code","source":"# Bert Parameters\nPRE_TRAINED_MODEL_NAME = 'bert-base-cased'\nMAX_LEN = 160\nBATCH_SIZE = 16\nEPOCHS = 1 # Can be increased but the training time will also increase significantly","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:05:04.174518Z","iopub.execute_input":"2025-02-04T21:05:04.174871Z","iopub.status.idle":"2025-02-04T21:05:04.178187Z","shell.execute_reply.started":"2025-02-04T21:05:04.174845Z","shell.execute_reply":"2025-02-04T21:05:04.177498Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Tokenizing","metadata":{}},{"cell_type":"code","source":"# Bert Tokenizer\ntokenizer = BertTokenizer.from_pretrained(PRE_TRAINED_MODEL_NAME)\n\n# Bert Tokenizer with example\nsample_txt = 'Corona Sucks so bad, my final year is kinda ruined, man!!'\n\ntokens = tokenizer.tokenize(sample_txt)\ntoken_ids = tokenizer.convert_tokens_to_ids(tokens)\nprint(f' Sentence: {sample_txt}')\nprint(f'   Tokens: {tokens}')\nprint(f'Token IDs: {token_ids}')\n\nencoding = tokenizer.encode_plus(\n    sample_txt,\n    max_length=32,\n    add_special_tokens=True,\n    return_token_type_ids=False,\n    padding='max_length',  \n    return_attention_mask=True,\n    return_tensors='pt',  # Return PyTorch tensors\n)\n\nencoding.keys()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:05:06.916341Z","iopub.execute_input":"2025-02-04T21:05:06.916657Z","iopub.status.idle":"2025-02-04T21:05:07.115125Z","shell.execute_reply.started":"2025-02-04T21:05:06.916633Z","shell.execute_reply":"2025-02-04T21:05:07.114236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Using the Bert tokenizer for encoding the questions\ntoken_lens = []\n\nfor txt in tqdm(df.processed_text):\n    encoding = tokenizer.encode_plus(\n        txt,\n        max_length=512,\n        truncation=True,\n        padding='max_length',  # Padding to max length\n        return_attention_mask=True,\n        return_token_type_ids=False,\n        return_tensors='pt'\n    )\n    token_lens.append(encoding['input_ids'].shape[1])  # Append the token length","metadata":{"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:05:11.032321Z","iopub.execute_input":"2025-02-04T21:05:11.032665Z","iopub.status.idle":"2025-02-04T21:16:44.585294Z","shell.execute_reply.started":"2025-02-04T21:05:11.032635Z","shell.execute_reply":"2025-02-04T21:16:44.584531Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Preparing the dataset with input ids and attention_masks\nclass GPquestionDataset(Dataset):\n    \n    def __init__(self, questions, targets, tokenizer, max_len):\n        self.questions = questions\n        self.targets = targets\n        self.tokenizer = tokenizer\n        self.max_len = max_len\n\n    def __len__(self):\n        return len(self.questions)\n\n    def __getitem__(self, item):\n        question = str(self.questions[item])\n        target = self.targets[item]\n\n        # Encoding the question text using the tokenizer\n        encoding = self.tokenizer.encode_plus(\n            question,\n            add_special_tokens=True,\n            max_length=self.max_len,\n            return_token_type_ids=False,\n            padding='max_length', \n            return_attention_mask=True,\n            return_tensors='pt',  # Return PyTorch tensors\n            truncation=True  \n        )\n\n        return {\n            'question_text': question,\n            'input_ids': encoding['input_ids'].squeeze(0),  # Remove the batch dimension (1,)\n            'attention_mask': encoding['attention_mask'].squeeze(0),  # Same here\n            'targets': torch.tensor(target, dtype=torch.long)\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:20:23.910840Z","iopub.execute_input":"2025-02-04T21:20:23.911277Z","iopub.status.idle":"2025-02-04T21:20:23.918642Z","shell.execute_reply.started":"2025-02-04T21:20:23.911242Z","shell.execute_reply":"2025-02-04T21:20:23.917846Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Splitting the Dataset and creating separate dataloaders","metadata":{}},{"cell_type":"code","source":"# Splitting the dataset for training, validation and testing\ndf_train, df_test = train_test_split(\n  df,\n  test_size = 0.4,\n  random_state = RANDOM_SEED\n)\ndf_val, df_test = train_test_split(\n  df_test,\n  test_size = 0.6,\n  random_state = RANDOM_SEED\n)\nprint (\"The shape of the training dataset : \", df_train.shape)\nprint (\"The shape of the validation dataset : \", df_val.shape)\nprint (\"The shape of the testing dataset : \", df_test.shape)","metadata":{"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:20:26.325548Z","iopub.execute_input":"2025-02-04T21:20:26.325883Z","iopub.status.idle":"2025-02-04T21:20:27.132057Z","shell.execute_reply.started":"2025-02-04T21:20:26.325856Z","shell.execute_reply":"2025-02-04T21:20:27.131309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Creating the data loader for training, validation, and testing using the PyTorch DataLoader\ndef create_data_loader(df, tokenizer, max_len, batch_size):\n    ds = GPquestionDataset(\n        questions = df.processed_text.to_numpy(),\n        targets = df.target.to_numpy(),\n        tokenizer = tokenizer,\n        max_len = max_len\n    )\n    return DataLoader(\n        ds,\n        batch_size = batch_size,\n        num_workers = 0\n    )\n\n# Creating data loaders for train, validation, and test sets\ntrain_data_loader = create_data_loader(df_train, tokenizer, MAX_LEN, BATCH_SIZE)\nval_data_loader = create_data_loader(df_val, tokenizer, MAX_LEN, BATCH_SIZE)\ntest_data_loader = create_data_loader(df_test, tokenizer, MAX_LEN, BATCH_SIZE)\n\n# Accessing one batch from the training data loader\ndata = next(iter(train_data_loader))\n\n# Checking the keys in the batch\nprint(data.keys())\n\n# Printing the number of batches in each DataLoader\nprint(len(train_data_loader))  # Number of batches in the training set\nprint(len(val_data_loader))    # Number of batches in the validation set\nprint(len(test_data_loader))   # Number of batches in the test set","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:21:13.076038Z","iopub.execute_input":"2025-02-04T21:21:13.076359Z","iopub.status.idle":"2025-02-04T21:21:13.142769Z","shell.execute_reply.started":"2025-02-04T21:21:13.076333Z","shell.execute_reply":"2025-02-04T21:21:13.142106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Shape of the torch\nprint(data['input_ids'].shape)\nprint(data['attention_mask'].shape)\nprint(data['targets'].shape)","metadata":{"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:21:16.618903Z","iopub.execute_input":"2025-02-04T21:21:16.619238Z","iopub.status.idle":"2025-02-04T21:21:16.623638Z","shell.execute_reply.started":"2025-02-04T21:21:16.619211Z","shell.execute_reply":"2025-02-04T21:21:16.622882Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Downloading Base BERT Model","metadata":{}},{"cell_type":"code","source":"# Using the Bert Model\nbert_model = BertModel.from_pretrained(PRE_TRAINED_MODEL_NAME)\n\n# Forward pass\noutput = bert_model(\n  input_ids=encoding['input_ids'],\n  attention_mask=encoding['attention_mask']\n)\n\n# Extracting the last hidden state correctly\nlast_hidden_state = output.last_hidden_state\npooled_output = output.pooler_output\n\nprint(last_hidden_state.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:21:26.398926Z","iopub.execute_input":"2025-02-04T21:21:26.399326Z","iopub.status.idle":"2025-02-04T21:21:27.511501Z","shell.execute_reply.started":"2025-02-04T21:21:26.399289Z","shell.execute_reply":"2025-02-04T21:21:27.510541Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Setting the Model Architecture","metadata":{}},{"cell_type":"code","source":"# Model with BERT layer and Dropout\nclass QuestionClassifier(nn.Module):\n    def __init__(self, n_classes):\n        super(QuestionClassifier, self).__init__()\n        # Loading pre-trained BERT model\n        self.bert = BertModel.from_pretrained(PRE_TRAINED_MODEL_NAME)\n        # Dropout layer with 30% probability of dropout\n        self.drop = nn.Dropout(p = 0.3)\n        # Final classification layer (Linear) with output size based on number of classes\n        self.out = nn.Linear(self.bert.config.hidden_size, n_classes)\n\n    def forward(self, input_ids, attention_mask):\n        # Forward pass through BERT model\n        outputs = self.bert(input_ids = input_ids, attention_mask = attention_mask)\n        # Pooled output\n        pooled_output = outputs[1]\n        # Apply dropout for regularization\n        output = self.drop(pooled_output)\n        # Final output layer (Linear)\n        return self.out(output)\n\n# Initialize model with 2 classes (for binary classification)\nmodel = QuestionClassifier(n_classes=2)\nmodel = model.to(device)  # Move the model to GPU in this case otherwise CPU","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:24:05.709000Z","iopub.execute_input":"2025-02-04T21:24:05.709340Z","iopub.status.idle":"2025-02-04T21:24:06.091258Z","shell.execute_reply.started":"2025-02-04T21:24:05.709315Z","shell.execute_reply":"2025-02-04T21:24:06.090540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"input_ids = data['input_ids'].to(device)\nattention_mask = data['attention_mask'].to(device)\n\nprint(input_ids.shape) # batch size x seq length\nprint(attention_mask.shape) # batch size x seq length\n\noutputs = model(input_ids = input_ids, attention_mask = attention_mask)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:24:41.496117Z","iopub.execute_input":"2025-02-04T21:24:41.496420Z","iopub.status.idle":"2025-02-04T21:24:41.526217Z","shell.execute_reply.started":"2025-02-04T21:24:41.496396Z","shell.execute_reply":"2025-02-04T21:24:41.525354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = AdamW(model.parameters(), lr = 2e-5, correct_bias = False)\ntotal_steps = len(train_data_loader) * EPOCHS\nscheduler = get_linear_schedule_with_warmup(\n  optimizer,\n  num_warmup_steps = 0,\n  num_training_steps = total_steps\n)\nloss_fn = nn.CrossEntropyLoss().to(device)","metadata":{"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:24:54.156232Z","iopub.execute_input":"2025-02-04T21:24:54.156576Z","iopub.status.idle":"2025-02-04T21:24:54.173197Z","shell.execute_reply.started":"2025-02-04T21:24:54.156549Z","shell.execute_reply":"2025-02-04T21:24:54.172359Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Training and Evaluation Functions","metadata":{}},{"cell_type":"code","source":"# Training function\ndef train_epoch(\n    model,\n    data_loader,\n    loss_fn,\n    optimizer,\n    device,\n    scheduler,\n    n_examples ): \n    \n    # Putting the model in training mode\n    model = model.train()\n    \n    losses = []\n    correct_predictions = 0\n    \n    for d in data_loader:\n        # Moving data to device (GPU or CPU)\n        input_ids = d[\"input_ids\"].to(device)\n        attention_mask = d[\"attention_mask\"].to(device)\n        targets = d[\"targets\"].to(device)\n        \n        # Forward pass\n        outputs = model(input_ids=input_ids, attention_mask=attention_mask)\n        \n        # Get predictions\n        _, preds = torch.max(outputs, dim=1)\n        \n        # Calculate loss\n        loss = loss_fn(outputs, targets)\n        \n        # Track correct predictions\n        correct_predictions += torch.sum(preds == targets).float()\n        \n        # Append loss for averaging later\n        losses.append(loss.item())\n        \n        # Backpropagation\n        loss.backward()\n        \n        # Gradient clipping to prevent exploding gradients\n        nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        \n        # Optimizer step\n        optimizer.step()\n        \n        # Scheduler step (for learning rate adjustment)\n        scheduler.step()\n        \n        # Zero gradients to prevent accumulation from previous steps\n        optimizer.zero_grad()\n        \n    # Return the average accuracy and loss for the epoch\n    accuracy = correct_predictions.double() / n_examples\n    avg_loss = np.mean(losses)\n    return accuracy, avg_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:27:47.948011Z","iopub.execute_input":"2025-02-04T21:27:47.948340Z","iopub.status.idle":"2025-02-04T21:27:47.953908Z","shell.execute_reply.started":"2025-02-04T21:27:47.948314Z","shell.execute_reply":"2025-02-04T21:27:47.953217Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Evaluation function\ndef eval_model(model, data_loader, loss_fn, device, n_examples):\n    \n    # Putting the model in evaluation mode\n    model = model.eval()\n    \n    losses = []\n    correct_predictions = 0\n    \n    with torch.no_grad():  # Disable gradient computation\n        for d in data_loader:\n            input_ids = d[\"input_ids\"].to(device)\n            attention_mask = d[\"attention_mask\"].to(device)\n            targets = d[\"targets\"].to(device)\n            \n            # Forward pass\n            outputs = model(input_ids=input_ids, attention_mask=attention_mask)\n            \n            # Get predictions\n            _, preds = torch.max(outputs, dim=1)\n            \n            # Calculate loss\n            loss = loss_fn(outputs, targets)\n            \n            # Track correct predictions\n            correct_predictions += torch.sum(preds == targets).float()\n            \n            # Append loss for averaging later\n            losses.append(loss.item())\n    \n    # Return accuracy and average loss\n    accuracy = correct_predictions.double() / n_examples\n    avg_loss = np.mean(losses)\n    return accuracy, avg_loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:28:51.057747Z","iopub.execute_input":"2025-02-04T21:28:51.058242Z","iopub.status.idle":"2025-02-04T21:28:51.063895Z","shell.execute_reply.started":"2025-02-04T21:28:51.058209Z","shell.execute_reply":"2025-02-04T21:28:51.062823Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Training and Evaluating the BERT Model","metadata":{}},{"cell_type":"code","source":"# The training loop\nstart_time = time.time() # To track the training time\nhistory = defaultdict(list)\nbest_accuracy = 0\n\nfor epoch in range(EPOCHS):\n    print(f'Epoch {epoch + 1}/{EPOCHS}')\n    print('-' * 70)\n    \n    # Training phase\n    train_acc, train_loss = train_epoch(\n        model,\n        train_data_loader,\n        loss_fn,\n        optimizer,\n        device,\n        scheduler,\n        len(df_train)\n    )\n    print(f'Train loss {train_loss:.4f} accuracy {train_acc:.4f}')\n    \n    # Validation phase\n    val_acc, val_loss = eval_model(\n        model,\n        val_data_loader,\n        loss_fn,\n        device,\n        len(df_val)\n    )\n    print(f'Val loss {val_loss:.4f} accuracy {val_acc:.4f}')\n    \n    # Log history\n    history['train_acc'].append(train_acc)\n    history['train_loss'].append(train_loss)\n    history['val_acc'].append(val_acc)\n    history['val_loss'].append(val_loss)\n    \n    # Checking if the model's validation accuracy is the best we've seen\n    if val_acc > best_accuracy:\n        print(f'Saving model with val_acc {val_acc:.4f}')\n        torch.save(model.state_dict(), 'best_model_state.bin')\n        best_accuracy = val_acc\n\nend_time = time.time()\nexecution_time = end_time - start_time\n\nprint(f\"Total execution time: {execution_time:.2f} seconds\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:43:33.946025Z","iopub.execute_input":"2025-02-04T21:43:33.946357Z","execution_failed":"2025-02-04T21:45:30.934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Making sure the model is in the right device for evaluation\nmodel = model.to(device) \n\n# Evaluating the model\ntest_acc, _ = eval_model(\n  model,\n  test_data_loader,\n  loss_fn,\n  device,\n  len(df_test)\n)\n\nprint (test_acc)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Predictions function\ndef get_predictions(model, data_loader):\n    model = model.eval()  # Set model to evaluation mode\n    question_texts = []\n    predictions = []\n    prediction_probs = []\n    real_values = []\n    \n    with torch.no_grad():  # Disable gradients for inference\n        for d in data_loader:\n            texts = d[\"question_text\"]\n            input_ids = d[\"input_ids\"].to(device)\n            attention_mask = d[\"attention_mask\"].to(device)\n            targets = d[\"targets\"].to(device)\n            \n            # Forward pass\n            outputs = model(input_ids=input_ids, attention_mask=attention_mask)\n            \n            # Gettting predictions (indices of max logits)\n            _, preds = torch.max(outputs, dim=1)\n            \n            # Collecting data\n            question_texts.extend(texts)\n            predictions.extend(preds)\n            prediction_probs.extend(torch.softmax(outputs, dim=1).cpu())  # Converting logits to probabilities\n            real_values.extend(targets)\n    \n    # Converting lists to tensors and move to CPU\n    predictions = torch.stack(predictions).cpu()\n    prediction_probs = torch.stack(prediction_probs).cpu()\n    real_values = torch.stack(real_values).cpu()\n    \n    return question_texts, predictions, prediction_probs, real_values\n\ny_question_texts, y_pred, y_pred_probs, y_test = get_predictions(\n    model,\n    test_data_loader\n)\n\n# Checking the results\nprint(\"Sample predictions:\", y_pred[:5])\nprint(\"Sample probabilities:\", y_pred_probs[:5])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# First 10 predictions and corresponding texts\ni = 0\nfor t, pred, prob in zip(y_question_texts, y_pred, y_pred_probs):\n    print(f\"Text: {t}\")\n    print(f\"Prediction: {pred}   Probabilities: {prob.numpy()}\")\n    print('-' * 50)  # Divider for readability\n    i += 1\n    if i == 10:\n        break","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Classification_report\nprint(classification_report(y_test, y_pred))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T21:36:17.865080Z","iopub.execute_input":"2025-02-04T21:36:17.865414Z","iopub.status.idle":"2025-02-04T21:36:17.889186Z","shell.execute_reply.started":"2025-02-04T21:36:17.865381Z","shell.execute_reply":"2025-02-04T21:36:17.888158Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Confusion matrix\nconf = confusion_matrix(y_test, y_pred)\nsns.heatmap(conf, annot=True, fmt=\"d\", cmap=\"Blues\", xticklabels=[\"Class 0\", \"Class 1\"], yticklabels=[\"Class 0\", \"Class 1\"])\nplt.ylabel(\"Actual\")\nplt.xlabel(\"Predicted\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Accuracy won't be useful since the dataset is highly skewed as shown in the EDA, so using f1_score for evaluation","metadata":{}},{"cell_type":"code","source":"# f1_score\nprint (f1_score(y_test, y_pred))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}