{"metadata":{"colab":{"provenance":[],"gpuType":"T4"},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"accelerator":"GPU","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":"gpu","dataSources":[{"sourceId":3043,"databundleVersionId":46668,"sourceType":"competition"},{"sourceId":2644,"sourceType":"modelInstanceVersion","modelInstanceId":1910},{"sourceId":2938,"sourceType":"modelInstanceVersion","modelInstanceId":2180}],"dockerImageVersionId":30683,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<h2 align=center> <b><u> Fine-Tuning BERT to predict Closed Questions on Stack Overflow </u></b>\n</h2>","metadata":{"id":"zGCJYkQj_Uu2"}},{"cell_type":"markdown","source":"### 1. Check GPU Availability and install dependencies","metadata":{"id":"mpe6GhLuBJWB"}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"id":"8V9c8vzSL3aj","outputId":"cc54afaa-d2fa-429f-8228-863c8423d1b6","execution":{"iopub.status.busy":"2024-05-07T02:21:06.796685Z","iopub.execute_input":"2024-05-07T02:21:06.797091Z","iopub.status.idle":"2024-05-07T02:21:07.811792Z","shell.execute_reply.started":"2024-05-07T02:21:06.797043Z","shell.execute_reply":"2024-05-07T02:21:07.810787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install tensorflow_text\n# After running this cell, we have to restart the Kernel!","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"scrolled":true,"execution":{"iopub.status.busy":"2024-05-07T02:21:41.985717Z","iopub.execute_input":"2024-05-07T02:21:41.986077Z","iopub.status.idle":"2024-05-07T02:21:54.517733Z","shell.execute_reply.started":"2024-05-07T02:21:41.986046Z","shell.execute_reply":"2024-05-07T02:21:54.516644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow_text as text  # Registers the ops.","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:21:54.519714Z","iopub.execute_input":"2024-05-07T02:21:54.520033Z","iopub.status.idle":"2024-05-07T02:22:07.471321Z","shell.execute_reply.started":"2024-05-07T02:21:54.519999Z","shell.execute_reply":"2024-05-07T02:22:07.470538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 2. Import the Closed Questions on Stack Overflow Dataset","metadata":{"id":"IMsEoT3Fg4Wg"}},{"cell_type":"code","source":"import numpy as np\nimport tensorflow as tf\nimport tensorflow_hub as hub","metadata":{"id":"GmqEylyFYTdP","outputId":"8e9f0646-e8d5-4279-9da5-2e66298f6762","execution":{"iopub.status.busy":"2024-05-07T02:22:07.472416Z","iopub.execute_input":"2024-05-07T02:22:07.472917Z","iopub.status.idle":"2024-05-07T02:22:07.477330Z","shell.execute_reply.started":"2024-05-07T02:22:07.472891Z","shell.execute_reply":"2024-05-07T02:22:07.476290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"TF Version: \", tf.__version__)\nprint(\"Eager mode: \", tf.executing_eagerly())\nprint(\"Hub version: \", hub.__version__)\nprint(\"GPU is\", \"available\" if tf.config.experimental.list_physical_devices(\"GPU\") else \"NOT AVAILABLE\")","metadata":{"id":"ZuX1lB8pPJ-W","outputId":"9d0eb66e-32c2-480f-ac6d-8f97b92838e3","execution":{"iopub.status.busy":"2024-05-07T02:22:07.478560Z","iopub.execute_input":"2024-05-07T02:22:07.478898Z","iopub.status.idle":"2024-05-07T02:22:07.677834Z","shell.execute_reply.started":"2024-05-07T02:22:07.478867Z","shell.execute_reply":"2024-05-07T02:22:07.676845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\n\ndf = pd.read_csv(\"/kaggle/input/predict-closed-questions-on-stack-overflow/train-sample.csv\")\ndf.shape","metadata":{"id":"0nI-9itVwCCQ","outputId":"fec5f1a4-561e-4302-ae08-aa1e1411165b","execution":{"iopub.status.busy":"2024-05-07T02:22:07.680686Z","iopub.execute_input":"2024-05-07T02:22:07.680999Z","iopub.status.idle":"2024-05-07T02:22:12.363969Z","shell.execute_reply.started":"2024-05-07T02:22:07.680972Z","shell.execute_reply":"2024-05-07T02:22:12.363084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.tail(5)","metadata":{"id":"yeHE98KiMvDd","outputId":"ffbfa52a-b91b-4e05-b223-ec88c8156c0c","execution":{"iopub.status.busy":"2024-05-07T02:22:12.365349Z","iopub.execute_input":"2024-05-07T02:22:12.365724Z","iopub.status.idle":"2024-05-07T02:22:12.387840Z","shell.execute_reply.started":"2024-05-07T02:22:12.365690Z","shell.execute_reply":"2024-05-07T02:22:12.386973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Assuming df is your DataFrame containing the 'OpenStatus' column\ndf['OpenStatus'].value_counts().plot(kind='bar', title='Class Distribution')\nplt.xlabel('OpenStatus')\nplt.ylabel('Count')\nplt.show()","metadata":{"id":"leRFRWJMocVa","outputId":"531296b7-0bc2-418d-ac01-461af40e8510","execution":{"iopub.status.busy":"2024-05-07T02:22:12.389023Z","iopub.execute_input":"2024-05-07T02:22:12.389340Z","iopub.status.idle":"2024-05-07T02:22:12.719909Z","shell.execute_reply.started":"2024-05-07T02:22:12.389273Z","shell.execute_reply":"2024-05-07T02:22:12.719006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Checking the number of samples from each class\n\nclass_distribution = df['OpenStatus'].value_counts()\nprint(class_distribution)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:12.721161Z","iopub.execute_input":"2024-05-07T02:22:12.721540Z","iopub.status.idle":"2024-05-07T02:22:12.748030Z","shell.execute_reply.started":"2024-05-07T02:22:12.721508Z","shell.execute_reply":"2024-05-07T02:22:12.746877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Checking for missing values\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport pandas as pd\n\n# Assuming 'df' is your DataFrame\n# Create a boolean DataFrame indicating missing values\nmissing_values = df.isnull()\n\n# Plot heatmap\nplt.figure(figsize=(10, 6))\nsns.heatmap(missing_values, cbar=False, cmap='viridis')\nplt.title('Missing Values Heatmap')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:12.749331Z","iopub.execute_input":"2024-05-07T02:22:12.749627Z","iopub.status.idle":"2024-05-07T02:22:15.372449Z","shell.execute_reply.started":"2024-05-07T02:22:12.749601Z","shell.execute_reply":"2024-05-07T02:22:15.371508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 3. Preparing Input Data for Training and Evaluation","metadata":{"id":"ELjswHcFHfp3"}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n\n# Splitting the dataset into train, validation and test (70%, 20% and 10% respectively)\ntrain_df, remaining = train_test_split(df, random_state=42, train_size=0.8, stratify=df.OpenStatus.values)\nvalid_df, test_df = train_test_split(remaining, random_state=42, train_size=0.50, stratify=remaining.OpenStatus.values)\n\n\n\n# Display the shapes of the downsampled training and validation datasets\ntrain_df.shape, valid_df.shape, test_df.shape","metadata":{"id":"fScULIGPwuWk","outputId":"f3f0dc5f-c7a9-4901-b3fe-dcff630b0eaa","execution":{"iopub.status.busy":"2024-05-07T02:22:15.373684Z","iopub.execute_input":"2024-05-07T02:22:15.374226Z","iopub.status.idle":"2024-05-07T02:22:15.723624Z","shell.execute_reply.started":"2024-05-07T02:22:15.374199Z","shell.execute_reply":"2024-05-07T02:22:15.722697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Checking the number of samples from each class of train_df\n\nclass_distribution = train_df['OpenStatus'].value_counts()\nprint(class_distribution)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:15.725958Z","iopub.execute_input":"2024-05-07T02:22:15.726255Z","iopub.status.idle":"2024-05-07T02:22:15.748540Z","shell.execute_reply.started":"2024-05-07T02:22:15.726230Z","shell.execute_reply":"2024-05-07T02:22:15.747684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The given dataset has a tremendous data imbalance for different classes. One possible approach can be to downsample all of 'open', 'not a real question', 'off topic' and 'not constructive' classes to the size of the 'too localized' class.","metadata":{}},{"cell_type":"code","source":"from sklearn.utils import resample\n\n# Separating different classes in the dataset\nopen_class = train_df[train_df.OpenStatus == 'open']\nnot_a_real_question_class = train_df[train_df.OpenStatus == 'not a real question']\noff_topic_class = train_df[train_df.OpenStatus == 'off topic']\nnot_constructive_class = train_df[train_df.OpenStatus == 'not constructive']\ntoo_localized_class = train_df[train_df.OpenStatus == 'too localized']\n\n\n# Downsampling the 5 classes to the size of the 'off topic' class\nopen_class = resample(open_class,\n                      replace=False,  # sample without replacement\n                      n_samples=len(too_localized_class),  # match target class size\n                      random_state=42)  # for reproducible results\n\nnot_a_real_question_class = resample(not_a_real_question_class,\n                                     replace=False,  # sample without replacement\n                                     n_samples=len(too_localized_class),  # match target class size\n                                     random_state=42)  # for reproducible results\n\noff_topic_class = resample(off_topic_class,\n                                     replace=False,  # sample without replacement\n                                     n_samples=len(too_localized_class),  # match target class size\n                                     random_state=42)  # for reproducible results\n\nnot_constructive_class = resample(not_constructive_class,\n                                     replace=False,  # sample without replacement\n                                     n_samples=len(too_localized_class),  # match target class size\n                                     random_state=42)  # for reproducible results\n\n# Combining all the minority class with the resampled classes\ntrain_df = pd.concat([open_class, \n                not_a_real_question_class, \n                off_topic_class, \n                not_constructive_class, \n                too_localized_class])\n\n# Shuffle the downsampled training dataset\ntrain_df = train_df.sample(frac=1, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:15.749801Z","iopub.execute_input":"2024-05-07T02:22:15.750092Z","iopub.status.idle":"2024-05-07T02:22:15.968694Z","shell.execute_reply.started":"2024-05-07T02:22:15.750061Z","shell.execute_reply":"2024-05-07T02:22:15.967880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.shape","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:15.970110Z","iopub.execute_input":"2024-05-07T02:22:15.970411Z","iopub.status.idle":"2024-05-07T02:22:15.975967Z","shell.execute_reply.started":"2024-05-07T02:22:15.970381Z","shell.execute_reply":"2024-05-07T02:22:15.975009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Checking the number of samples from each class of train_df\n\nclass_distribution = train_df['OpenStatus'].value_counts()\nprint(class_distribution)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:15.980746Z","iopub.execute_input":"2024-05-07T02:22:15.981602Z","iopub.status.idle":"2024-05-07T02:22:15.991330Z","shell.execute_reply.started":"2024-05-07T02:22:15.981570Z","shell.execute_reply":"2024-05-07T02:22:15.990240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.tail(5)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:15.992460Z","iopub.execute_input":"2024-05-07T02:22:15.992714Z","iopub.status.idle":"2024-05-07T02:22:16.014331Z","shell.execute_reply.started":"2024-05-07T02:22:15.992692Z","shell.execute_reply":"2024-05-07T02:22:16.013322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I have decided to use only the \"Title\", \"BodyMarkdown\" and \"OpenStatus\" columns from the dataframe. The \"Title\" and  \"BodyMarkdown\" can be joined together to create the text input and the \"OpenStatus\" column is our target column.","metadata":{}},{"cell_type":"code","source":"selected_columns = ['Title', 'BodyMarkdown', 'OpenStatus', 'Tag1', 'Tag2', 'Tag3', 'Tag4', 'Tag5']\ntrain_df = train_df[selected_columns]\nvalid_df = valid_df[selected_columns]\ntest_df = test_df[selected_columns]","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:16.015508Z","iopub.execute_input":"2024-05-07T02:22:16.015840Z","iopub.status.idle":"2024-05-07T02:22:16.035149Z","shell.execute_reply.started":"2024-05-07T02:22:16.015807Z","shell.execute_reply":"2024-05-07T02:22:16.034183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.tail(5)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:16.036391Z","iopub.execute_input":"2024-05-07T02:22:16.036779Z","iopub.status.idle":"2024-05-07T02:22:16.049744Z","shell.execute_reply.started":"2024-05-07T02:22:16.036742Z","shell.execute_reply":"2024-05-07T02:22:16.048711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here, I have combined the \"Title\", \"Tag1 to Tag5\" and \"BodyMarkdown\" columns to create a single text input.","metadata":{}},{"cell_type":"code","source":"for index, row in train_df.iterrows():\n    text = \"Title: \" + \"'\" + row.Title + \"'\" \n    #text += \" Tags: {\" + ', '.join(str(tag) for tag in [row.Tag1, row.Tag2, row.Tag3, row.Tag4, row.Tag5] if not pd.isnull(tag)) + \"}\"\n    text += \"  Body: \" + \"'\" + row.BodyMarkdown + \"'\"\n    train_df.at[index, 'text'] = text\n    \n\nfor index, row in valid_df.iterrows():\n    text = \"Title: \" + \"'\" + row.Title + \"'\" \n    #text += \" Tags: {\" + ', '.join(str(tag) for tag in [row.Tag1, row.Tag2, row.Tag3, row.Tag4, row.Tag5] if not pd.isnull(tag)) + \"}\" \n    text += \"  Body: \" + \"'\" + row.BodyMarkdown + \"'\"\n    valid_df.at[index, 'text'] = text\n\nfor index, row in test_df.iterrows():\n    text = \"Title: \" + \"'\" + row.Title + \"'\" \n    #text += \" Tags: {\" + ', '.join(str(tag) for tag in [row.Tag1, row.Tag2, row.Tag3, row.Tag4, row.Tag5] if not pd.isnull(tag)) + \"}\"\n    text += \"  Body: \" + \"'\" + row.BodyMarkdown + \"'\"\n    test_df.at[index, 'text'] = text","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:16.051241Z","iopub.execute_input":"2024-05-07T02:22:16.051967Z","iopub.status.idle":"2024-05-07T02:22:20.976531Z","shell.execute_reply.started":"2024-05-07T02:22:16.051938Z","shell.execute_reply":"2024-05-07T02:22:20.975704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.tail(5)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:20.977671Z","iopub.execute_input":"2024-05-07T02:22:20.977970Z","iopub.status.idle":"2024-05-07T02:22:20.991866Z","shell.execute_reply.started":"2024-05-07T02:22:20.977945Z","shell.execute_reply":"2024-05-07T02:22:20.990945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now, I have dropped the \"Title\" and \"BodyMarkdown\" columns from the dataframe.","metadata":{}},{"cell_type":"code","source":"columns_to_drop = ['Title', 'BodyMarkdown']\ntrain_df.drop(columns=columns_to_drop, inplace=True)\nvalid_df.drop(columns=columns_to_drop, inplace=True)\ntest_df.drop(columns=columns_to_drop, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:20.993150Z","iopub.execute_input":"2024-05-07T02:22:20.994149Z","iopub.status.idle":"2024-05-07T02:22:21.011244Z","shell.execute_reply.started":"2024-05-07T02:22:20.994107Z","shell.execute_reply":"2024-05-07T02:22:21.010444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.shape","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:21.012529Z","iopub.execute_input":"2024-05-07T02:22:21.013436Z","iopub.status.idle":"2024-05-07T02:22:21.018946Z","shell.execute_reply.started":"2024-05-07T02:22:21.013409Z","shell.execute_reply":"2024-05-07T02:22:21.018045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.tail(5)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:21.019896Z","iopub.execute_input":"2024-05-07T02:22:21.020148Z","iopub.status.idle":"2024-05-07T02:22:21.036466Z","shell.execute_reply.started":"2024-05-07T02:22:21.020126Z","shell.execute_reply":"2024-05-07T02:22:21.035639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for index, row in train_df.iterrows():\n    print(\"Text: \")\n    print(\"________________________\")\n    print(row.text, end='\\n\\n\\n\\n')\n    print(\"Target Class:\")\n    print(\"________________________\")\n    print(row.OpenStatus)\n    break","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:21.037694Z","iopub.execute_input":"2024-05-07T02:22:21.038442Z","iopub.status.idle":"2024-05-07T02:22:21.053586Z","shell.execute_reply.started":"2024-05-07T02:22:21.038414Z","shell.execute_reply":"2024-05-07T02:22:21.052598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we have to encode the Target column into numerical values.","metadata":{}},{"cell_type":"code","source":"# Definelabel mapping\ncustom_label_mapping = {'open': 0, \n                        'not a real question': 1, \n                        'off topic': 2,\n                        'not constructive': 3, \n                        'too localized': 4\n                       }  \n\n# Initialize LabelEncoder with custom mapping\nencoded_labels = []\nfor index, row in train_df.iterrows():\n    label = row['OpenStatus']\n    encoded_label = custom_label_mapping[label]\n    encoded_labels.append(encoded_label)\n\n# Add the encoded labels as a new column in the DataFrame\ntrain_df['OpenStatus_encoded'] = encoded_labels\n\n\n\n# Applying the same changes to valid_df and test_df\n\nencoded_labels = []\nfor index, row in valid_df.iterrows():\n    label = row['OpenStatus']\n    encoded_label = custom_label_mapping[label]\n    encoded_labels.append(encoded_label)\n\n# Add the encoded labels as a new column in the DataFrame\nvalid_df['OpenStatus_encoded'] = encoded_labels\n\n\n\nencoded_labels = []\nfor index, row in test_df.iterrows():\n    label = row['OpenStatus']\n    encoded_label = custom_label_mapping[label]\n    encoded_labels.append(encoded_label)\n\n# Add the encoded labels as a new column in the DataFrame\ntest_df['OpenStatus_encoded'] = encoded_labels","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:21.054726Z","iopub.execute_input":"2024-05-07T02:22:21.054983Z","iopub.status.idle":"2024-05-07T02:22:23.960170Z","shell.execute_reply.started":"2024-05-07T02:22:21.054961Z","shell.execute_reply":"2024-05-07T02:22:23.959167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.tail(5)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:23.961253Z","iopub.execute_input":"2024-05-07T02:22:23.961587Z","iopub.status.idle":"2024-05-07T02:22:23.975491Z","shell.execute_reply.started":"2024-05-07T02:22:23.961562Z","shell.execute_reply":"2024-05-07T02:22:23.974441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.shape","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:23.976756Z","iopub.execute_input":"2024-05-07T02:22:23.977121Z","iopub.status.idle":"2024-05-07T02:22:23.987196Z","shell.execute_reply.started":"2024-05-07T02:22:23.977071Z","shell.execute_reply":"2024-05-07T02:22:23.986341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Assuming df is your DataFrame containing the 'OpenStatus' column\ntrain_df['OpenStatus'].value_counts().plot(kind='bar', title='Class Distribution')\nplt.xlabel('OpenStatus')\nplt.ylabel('Count')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:23.988536Z","iopub.execute_input":"2024-05-07T02:22:23.989118Z","iopub.status.idle":"2024-05-07T02:22:24.254703Z","shell.execute_reply.started":"2024-05-07T02:22:23.989087Z","shell.execute_reply":"2024-05-07T02:22:24.253764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train = train_df[\"text\"]\ny_train = train_df[\"OpenStatus_encoded\"]\n\nX_valid = valid_df[\"text\"]\ny_valid = valid_df[\"OpenStatus_encoded\"]\n\nX_test = test_df[\"text\"]\ny_test = test_df[\"OpenStatus_encoded\"]","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:24.256257Z","iopub.execute_input":"2024-05-07T02:22:24.256566Z","iopub.status.idle":"2024-05-07T02:22:24.261754Z","shell.execute_reply.started":"2024-05-07T02:22:24.256540Z","shell.execute_reply":"2024-05-07T02:22:24.260762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(X_train.shape)\nprint(y_train.shape)\n\nprint(X_valid.shape)\nprint(y_valid.shape)\n\nprint(X_test.shape)\nprint(y_test.shape)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:24.262930Z","iopub.execute_input":"2024-05-07T02:22:24.263274Z","iopub.status.idle":"2024-05-07T02:22:24.274064Z","shell.execute_reply.started":"2024-05-07T02:22:24.263242Z","shell.execute_reply":"2024-05-07T02:22:24.273124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 4. Input Format for BERT","metadata":{"id":"9QinzNq6OsP1"}},{"cell_type":"markdown","source":"**Token IDs** - This corresponds to the tokenized strings padded with 0s\n                upto the max sequence length and beginning with CLS and ending with SEP. <br><br>\n**Input Mask** - Note that BERT uses Self-Attention Networks to provide\n                 contextualised embeddings corresponding to each token in the token string i.e., for each word in the string BERT looks to the left and right of it in the sentence so as to find contextual meaning of the word in the sentence (say, if there is a \"the\", then look at the noun to which it points). Now, note that we have padded our token strings with 0s upto the max seq length, but we do not want the padding 0s to influence the contextual information to be derived.  The Input Mask is a list of same length as the length of Token Ids (ie the max seq length) where there is a 0 for a padding and 1 for a valid token. The 0s will cancel out the internal multiplications that we perform for capturing the Self Attention for contextual information. <br><br>\n\n**Input Type IDs** - Note that originally BERT was pretrained on two   tasks, Masked Language Modelling (where random words from the sentence would be masked and it would be the task for the BERT to predict what those masked words are) and the other task was Next Sentence Prediction or NSP (Given two sentences, the BERT has to predict which came first and which came after. The first sentence was given the value 0 and the next sentence was given the value 1).<br><br>\n**In Text classification, we are dealing with only 1 sequence at a time, so our input type IDs would just be a vector with all values 0. **\n","metadata":{"id":"shyvv_0JaIzj"}},{"cell_type":"markdown","source":"### 5. Checking the tokenization process","metadata":{}},{"cell_type":"code","source":"preprocessor = hub.KerasLayer(\n    \"https://kaggle.com/models/tensorflow/bert/frameworks/TensorFlow2/variations/en-uncased-preprocess/versions/3\")\n\n# Tokenize the input text\ninput_text = [\"Hello, how are you?\"]\ntokenized_output = preprocessor(input_text)\n\n# Print token IDs\nprint(tokenized_output['input_word_ids'])\nprint(tokenized_output['input_mask'])\nprint(tokenized_output['input_type_ids'])","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:24.275248Z","iopub.execute_input":"2024-05-07T02:22:24.275613Z","iopub.status.idle":"2024-05-07T02:22:28.345571Z","shell.execute_reply.started":"2024-05-07T02:22:24.275588Z","shell.execute_reply":"2024-05-07T02:22:28.344496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Upon checking the input word ids, we find that the tokens are --> \"Hello\", \"#,\", \"how\", \"are\", \"you\" and \"#?\" {'#' signifies that the succeeding character ',' is attached to characters before] and they are encoded as [7592 1010 2129 2024 2017 1029]. \n\nNote that this is not the whole story. We also have to make sure that each sequence is initiated with the CLS token (which signifies start of sequence and has a token_id 101) and ended with SEP (separator) which means end of sequence and has a token_id 102. Also we have to make sure that all tensors have sequence size equal to the max_sequence_length by using padding.","metadata":{}},{"cell_type":"markdown","source":"### 6. Add a Classification Head to the BERT Layer","metadata":{"id":"GZxe-7yhPyQe"}},{"cell_type":"markdown","source":"We only need the pooled_output that represents the whole sentence (using the CLS token that contains the contextual information of the whole sentence) and not the sequence_output.","metadata":{"id":"eED7TDu0vQ2Y"}},{"cell_type":"code","source":"# Building the model\n\ntext_input = tf.keras.layers.Input(shape=(), dtype=tf.string)\npreprocessor = hub.KerasLayer(\n    \"https://kaggle.com/models/tensorflow/bert/frameworks/TensorFlow2/variations/en-uncased-preprocess/versions/3\")\nencoder_inputs = preprocessor(text_input)\nencoder = hub.KerasLayer(\n    \"https://www.kaggle.com/models/tensorflow/bert/frameworks/TensorFlow2/variations/bert-en-uncased-l-12-h-768-a-12/versions/2\",\n    trainable=True)\noutputs = encoder(encoder_inputs)\npooled_output = outputs[\"pooled_output\"]      # [batch_size, 768].\nsequence_output = outputs[\"sequence_output\"]  # [batch_size, seq_length, 768].\n\n\n\n\n\n# Classification\n# Add dropout layer\ndrop1 = tf.keras.layers.Dropout(0.5)(pooled_output)\nbatch_norm1 = tf.keras.layers.BatchNormalization()(drop1)\n\n# Add hidden dense layers\nhidden1 = tf.keras.layers.Dense(512, activation='relu')(batch_norm1)\ndrop2 = tf.keras.layers.Dropout(0.4)(hidden1)\nbatch_norm2 = tf.keras.layers.BatchNormalization()(drop2)\nhidden2 = tf.keras.layers.Dense(128, activation='relu')(batch_norm2)\ndrop3 = tf.keras.layers.Dropout(0.3)(hidden2)\nbatch_norm3 = tf.keras.layers.BatchNormalization()(drop3)\nhidden3 = tf.keras.layers.Dense(32, activation='relu')(batch_norm3)\ndrop4 = tf.keras.layers.Dropout(0.2)(hidden3)\nbatch_norm4 = tf.keras.layers.BatchNormalization()(drop4)\n\n# Output layer\noutput_layer = tf.keras.layers.Dense(5, activation='softmax', name='output')(batch_norm4)\n\n\nmodel=tf.keras.Model(inputs=[text_input],outputs=[output_layer])","metadata":{"id":"G9il4gtlADcp","execution":{"iopub.status.busy":"2024-05-07T02:22:28.346704Z","iopub.execute_input":"2024-05-07T02:22:28.347040Z","iopub.status.idle":"2024-05-07T02:22:47.594516Z","shell.execute_reply.started":"2024-05-07T02:22:28.347012Z","shell.execute_reply":"2024-05-07T02:22:47.593576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 7. Fine-Tune BERT for Text Classification","metadata":{"id":"S6maM-vr7YaJ"}},{"cell_type":"code","source":"model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=2e-5),\n              loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),\n              metrics = ['accuracy'])\nmodel.summary()","metadata":{"id":"ptCtiiONsBgo","outputId":"5586d260-0153-4659-e47e-a5420023debb","execution":{"iopub.status.busy":"2024-05-07T02:22:47.595898Z","iopub.execute_input":"2024-05-07T02:22:47.596773Z","iopub.status.idle":"2024-05-07T02:22:47.684035Z","shell.execute_reply.started":"2024-05-07T02:22:47.596734Z","shell.execute_reply":"2024-05-07T02:22:47.683109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define callbacks\ncallbacks = [\n    tf.keras.callbacks.ModelCheckpoint(\n        filepath='best_model.h5',  # Path to save the best model\n        save_best_only=True,  # Save only the best model\n        monitor='val_accuracy',  # Quantity to be monitored\n        save_weights_only=True,  # Do not save the entire model\n        verbose=1,  # Verbosity mode. 0 or 1.\n        save_freq='epoch'  # Save the model at the end of every epoch\n    ),\n    tf.keras.callbacks.EarlyStopping(\n        patience=3,  # Number of epochs with no improvement after which training will be stopped\n        monitor='val_accuracy',  # Quantity to be monitored\n        restore_best_weights=True  # Restore model weights from the epoch with the best value of the monitored quantity\n    ),\n    tf.keras.callbacks.ReduceLROnPlateau(\n        monitor='val_accuracy',  # Quantity to be monitored\n        factor=0.5,  # Factor by which the learning rate will be reduced. new_lr = lr * factor\n        patience=3,  # Number of epochs with no improvement after which learning rate will be reduced\n        min_lr=1e-10  # Lower bound on the learning rate\n    )\n]","metadata":{"execution":{"iopub.status.busy":"2024-05-07T02:22:47.685143Z","iopub.execute_input":"2024-05-07T02:22:47.685436Z","iopub.status.idle":"2024-05-07T02:22:47.691611Z","shell.execute_reply.started":"2024-05-07T02:22:47.685410Z","shell.execute_reply":"2024-05-07T02:22:47.690680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot the model architecture\ntf.keras.utils.plot_model(model, show_shapes=True, dpi=76)","metadata":{"id":"6GJaFnkbMtPL","outputId":"83dad392-eb4c-481b-fd54-5955c91c1197","execution":{"iopub.status.busy":"2024-05-07T02:22:47.692886Z","iopub.execute_input":"2024-05-07T02:22:47.693185Z","iopub.status.idle":"2024-05-07T02:22:48.161492Z","shell.execute_reply.started":"2024-05-07T02:22:47.693162Z","shell.execute_reply":"2024-05-07T02:22:48.160437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train model\nepochs = 10\nhistory = model.fit(X_train, \n                    y_train,\n                    validation_data = (X_valid, y_valid),\n                    epochs=epochs,\n                    verbose=1,\n                    callbacks=callbacks\n                   )","metadata":{"id":"OcREcgPUHr9O","outputId":"c9fa3758-6c20-4158-d22c-18da8aa40dac","execution":{"iopub.status.busy":"2024-05-07T02:22:48.162898Z","iopub.execute_input":"2024-05-07T02:22:48.163244Z","iopub.status.idle":"2024-05-07T03:22:01.599724Z","shell.execute_reply.started":"2024-05-07T02:22:48.163213Z","shell.execute_reply":"2024-05-07T03:22:01.598761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 8. Check the loss and accuracy curves","metadata":{"id":"kNZl1lx_cA5Y"}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef plot_graphs(history, metric):\n    plt.plot(history.history[metric])\n    plt.plot(history.history['val_'+metric], '')\n    plt.xlabel(\"Epochs\")\n    plt.ylabel(metric)\n    plt.legend([metric, 'val_'+metric])\n    plt.show()","metadata":{"id":"dCjgrUYH_IsE","execution":{"iopub.status.busy":"2024-05-07T03:22:01.601200Z","iopub.execute_input":"2024-05-07T03:22:01.601803Z","iopub.status.idle":"2024-05-07T03:22:01.607677Z","shell.execute_reply.started":"2024-05-07T03:22:01.601766Z","shell.execute_reply":"2024-05-07T03:22:01.606639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_graphs(history, 'loss')","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:22:01.608735Z","iopub.execute_input":"2024-05-07T03:22:01.609005Z","iopub.status.idle":"2024-05-07T03:22:01.801402Z","shell.execute_reply.started":"2024-05-07T03:22:01.608982Z","shell.execute_reply":"2024-05-07T03:22:01.800449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_graphs(history, 'accuracy')","metadata":{"id":"opu9neBA_98R","outputId":"e1b12c79-ff53-4f5a-c8b6-92e0b52a5b54","execution":{"iopub.status.busy":"2024-05-07T03:22:01.802815Z","iopub.execute_input":"2024-05-07T03:22:01.803475Z","iopub.status.idle":"2024-05-07T03:22:02.062372Z","shell.execute_reply.started":"2024-05-07T03:22:01.803439Z","shell.execute_reply":"2024-05-07T03:22:02.061486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 9. Evaluating on Test Dataframe","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score, log_loss\n\n\n# Obtain predictions from the BERT model\npredictions = model.predict(X_test)\n\n# Convert predictions to class labels\npredicted_labels = tf.argmax(predictions, axis=1)\n\n# Compute accuracy\naccuracy = accuracy_score(y_test, predicted_labels)\nprint(\"Accuracy:\", round(accuracy * 100, 4) , \"%\")\n\n# Compute sparse categorical cross-entropy loss\nloss = log_loss(y_test, predictions)\nprint(\"Sparse Categorical Cross-Entropy Loss:\", loss)","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:22:02.063938Z","iopub.execute_input":"2024-05-07T03:22:02.064332Z","iopub.status.idle":"2024-05-07T03:23:36.009329Z","shell.execute_reply.started":"2024-05-07T03:22:02.064295Z","shell.execute_reply":"2024-05-07T03:23:36.008350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Confusion Matrix","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n# Compute confusion matrix\nconf_matrix = confusion_matrix(y_test, predicted_labels)\n\n# Plot confusion matrix\nplt.figure(figsize=(8, 6))\nsns.heatmap(conf_matrix, annot=True, fmt=\"d\", cmap=\"Blues\", \n            xticklabels=[str(i) for i in range(5)], \n            yticklabels=[str(i) for i in range(5)])\nplt.xlabel('Predicted labels')\nplt.ylabel('True labels')\nplt.title('Confusion Matrix')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:23:36.010737Z","iopub.execute_input":"2024-05-07T03:23:36.011117Z","iopub.status.idle":"2024-05-07T03:23:36.359961Z","shell.execute_reply.started":"2024-05-07T03:23:36.011083Z","shell.execute_reply":"2024-05-07T03:23:36.359050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 10. Sanity Checks","metadata":{}},{"cell_type":"code","source":"import random\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:23:36.366024Z","iopub.execute_input":"2024-05-07T03:23:36.366349Z","iopub.status.idle":"2024-05-07T03:23:36.370504Z","shell.execute_reply.started":"2024-05-07T03:23:36.366322Z","shell.execute_reply":"2024-05-07T03:23:36.369463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_test_list = X_test.tolist()\ny_test_list = y_test.tolist()","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:23:36.371634Z","iopub.execute_input":"2024-05-07T03:23:36.371959Z","iopub.status.idle":"2024-05-07T03:23:36.383934Z","shell.execute_reply.started":"2024-05-07T03:23:36.371933Z","shell.execute_reply":"2024-05-07T03:23:36.383146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(X_test_list))\nprint(len(y_test_list))","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:23:36.385025Z","iopub.execute_input":"2024-05-07T03:23:36.385922Z","iopub.status.idle":"2024-05-07T03:23:36.395663Z","shell.execute_reply.started":"2024-05-07T03:23:36.385889Z","shell.execute_reply":"2024-05-07T03:23:36.394779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"custom_label_mapping = <br>{<br>'open': 0, <br>\n                        'not a real question': 1, <br>\n                        'off topic': 2, <br>\n                        'not constructive': 3, <br>\n                        'too localized': 4 <br>\n                       }  ","metadata":{}},{"cell_type":"code","source":"label_to_class_map = {0 : 'open',\n                      1 : 'not a real question',\n                      2 : 'off topic',\n                      3 : 'not constructive',\n                      4 : 'too localized'}","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:23:36.396669Z","iopub.execute_input":"2024-05-07T03:23:36.396936Z","iopub.status.idle":"2024-05-07T03:23:36.405916Z","shell.execute_reply.started":"2024-05-07T03:23:36.396913Z","shell.execute_reply":"2024-05-07T03:23:36.405176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Check 1","metadata":{}},{"cell_type":"code","source":"# Generate a random integer between 1 and 100 (inclusive)\nrand_int = random.randint(0, 14028)\n\nprint(\"___________________________________________________\")\nprint(\"\\n\\nTEXT:\\n\\n\")\nprint(X_test_list[rand_int])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\nprint(\"\\n\\nGROUND TRUTH:\\n\\n\")\nprint(label_to_class_map[y_test_list[rand_int]])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\nprint(\"\\n\\n\")\nprediction = model.predict([X_test_list[rand_int]])\nprint(\"\\n\\nPREDICTION:\\n\\n\")\n\ncategories = ['open', \n              'not a real question', \n              'off topic',\n              'not constructive',\n              'too localized'\n             ]\n\nplt.figure(figsize=(8, 6))\nplt.bar(categories, prediction[0])\nplt.xlabel('Classes')\nplt.ylabel('Probabilities')\nplt.title('Predictions')\nplt.show()\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\n\n\n\npredicted_class_index = np.argmax(prediction)\nprint(\"\\n\\nPREDICTED CLASS: \", label_to_class_map[predicted_class_index])\nprint('\\n\\n')\nprint(\"___________________________________________________\")\n\n\n\nprint(\"\\n\\n\")\nif label_to_class_map[predicted_class_index] == label_to_class_map[y_test_list[rand_int]]:\n    print(\"CORRECT PREDICTION !!!\")\nelse:\n    print(\"WRONG PREDICTION !!!\")","metadata":{"execution":{"iopub.status.busy":"2024-05-07T04:03:11.100268Z","iopub.execute_input":"2024-05-07T04:03:11.101202Z","iopub.status.idle":"2024-05-07T04:03:11.470814Z","shell.execute_reply.started":"2024-05-07T04:03:11.101167Z","shell.execute_reply":"2024-05-07T04:03:11.469853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Check 2","metadata":{}},{"cell_type":"code","source":"# Generate a random integer between 1 and 100 (inclusive)\nrand_int = random.randint(0, 14028)\n\nprint(\"___________________________________________________\")\nprint(\"\\n\\nTEXT:\\n\\n\")\nprint(X_test_list[rand_int])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\nprint(\"\\n\\nGROUND TRUTH:\\n\\n\")\nprint(label_to_class_map[y_test_list[rand_int]])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\nprint(\"\\n\\n\")\nprediction = model.predict([X_test_list[rand_int]])\nprint(\"\\n\\nPREDICTION:\\n\\n\")\n\ncategories = ['open', \n              'not a real question', \n              'off topic',\n              'not constructive',\n              'too localized'\n             ]\n\nplt.figure(figsize=(8, 6))\nplt.bar(categories, prediction[0])\nplt.xlabel('Classes')\nplt.ylabel('Probabilities')\nplt.title('Predictions')\nplt.show()\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\n\n\n\npredicted_class_index = np.argmax(prediction)\nprint(\"\\n\\nPREDICTED CLASS: \", label_to_class_map[predicted_class_index])\nprint('\\n\\n')\nprint(\"___________________________________________________\")\n\n\n\nprint(\"\\n\\n\")\nif label_to_class_map[predicted_class_index] == label_to_class_map[y_test_list[rand_int]]:\n    print(\"CORRECT PREDICTION !!!\")\nelse:\n    print(\"WRONG PREDICTION !!!\")","metadata":{"execution":{"iopub.status.busy":"2024-05-07T04:02:56.391132Z","iopub.execute_input":"2024-05-07T04:02:56.391534Z","iopub.status.idle":"2024-05-07T04:02:56.686508Z","shell.execute_reply.started":"2024-05-07T04:02:56.391499Z","shell.execute_reply":"2024-05-07T04:02:56.685535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Check 3","metadata":{}},{"cell_type":"code","source":"# Generate a random integer between 1 and 100 (inclusive)\nrand_int = random.randint(0, 14028)\n\nprint(\"___________________________________________________\")\nprint(\"\\n\\nTEXT:\\n\\n\")\nprint(X_test_list[rand_int])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\nprint(\"\\n\\nGROUND TRUTH:\\n\\n\")\nprint(label_to_class_map[y_test_list[rand_int]])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\nprint(\"\\n\\n\")\nprediction = model.predict([X_test_list[rand_int]])\nprint(\"\\n\\nPREDICTION:\\n\\n\")\n\ncategories = ['open', \n              'not a real question', \n              'off topic',\n              'not constructive',\n              'too localized'\n             ]\n\nplt.figure(figsize=(8, 6))\nplt.bar(categories, prediction[0])\nplt.xlabel('Classes')\nplt.ylabel('Probabilities')\nplt.title('Predictions')\nplt.show()\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\n\n\n\npredicted_class_index = np.argmax(prediction)\nprint(\"\\n\\nPREDICTED CLASS: \", label_to_class_map[predicted_class_index])\nprint('\\n\\n')\nprint(\"___________________________________________________\")\n\n\n\nprint(\"\\n\\n\")\nif label_to_class_map[predicted_class_index] == label_to_class_map[y_test_list[rand_int]]:\n    print(\"CORRECT PREDICTION !!!\")\nelse:\n    print(\"WRONG PREDICTION !!!\")","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:23:37.850696Z","iopub.execute_input":"2024-05-07T03:23:37.850998Z","iopub.status.idle":"2024-05-07T03:23:38.201812Z","shell.execute_reply.started":"2024-05-07T03:23:37.850965Z","shell.execute_reply":"2024-05-07T03:23:38.200854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Check 4","metadata":{}},{"cell_type":"code","source":"# Generate a random integer between 1 and 100 (inclusive)\nrand_int = random.randint(0, 14028)\n\nprint(\"___________________________________________________\")\nprint(\"\\n\\nTEXT:\\n\\n\")\nprint(X_test_list[rand_int])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\nprint(\"\\n\\nGROUND TRUTH:\\n\\n\")\nprint(label_to_class_map[y_test_list[rand_int]])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\nprint(\"\\n\\n\")\nprediction = model.predict([X_test_list[rand_int]])\nprint(\"\\n\\nPREDICTION:\\n\\n\")\n\ncategories = ['open', \n              'not a real question', \n              'off topic',\n              'not constructive',\n              'too localized'\n             ]\n\nplt.figure(figsize=(8, 6))\nplt.bar(categories, prediction[0])\nplt.xlabel('Classes')\nplt.ylabel('Probabilities')\nplt.title('Predictions')\nplt.show()\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\n\n\n\npredicted_class_index = np.argmax(prediction)\nprint(\"\\n\\nPREDICTED CLASS: \", label_to_class_map[predicted_class_index])\nprint('\\n\\n')\nprint(\"___________________________________________________\")\n\n\n\nprint(\"\\n\\n\")\nif label_to_class_map[predicted_class_index] == label_to_class_map[y_test_list[rand_int]]:\n    print(\"CORRECT PREDICTION !!!\")\nelse:\n    print(\"WRONG PREDICTION !!!\")","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:23:38.203009Z","iopub.execute_input":"2024-05-07T03:23:38.203333Z","iopub.status.idle":"2024-05-07T03:23:38.546321Z","shell.execute_reply.started":"2024-05-07T03:23:38.203298Z","shell.execute_reply":"2024-05-07T03:23:38.545290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Check 5","metadata":{}},{"cell_type":"code","source":"# Generate a random integer between 1 and 100 (inclusive)\nrand_int = random.randint(0, 14028)\n\nprint(\"___________________________________________________\")\nprint(\"\\n\\nTEXT:\\n\\n\")\nprint(X_test_list[rand_int])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\nprint(\"\\n\\nGROUND TRUTH:\\n\\n\")\nprint(label_to_class_map[y_test_list[rand_int]])\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\nprint(\"\\n\\n\")\nprediction = model.predict([X_test_list[rand_int]])\nprint(\"\\n\\nPREDICTION:\\n\\n\")\n\ncategories = ['open', \n              'not a real question', \n              'off topic',\n              'not constructive',\n              'too localized'\n             ]\n\nplt.figure(figsize=(8, 6))\nplt.bar(categories, prediction[0])\nplt.xlabel('Classes')\nplt.ylabel('Probabilities')\nplt.title('Predictions')\nplt.show()\nprint(\"\\n\\n\")\nprint(\"___________________________________________________\")\n\n\n\n\n\npredicted_class_index = np.argmax(prediction)\nprint(\"\\n\\nPREDICTED CLASS: \", label_to_class_map[predicted_class_index])\nprint('\\n\\n')\nprint(\"___________________________________________________\")\n\n\n\nprint(\"\\n\\n\")\nif label_to_class_map[predicted_class_index] == label_to_class_map[y_test_list[rand_int]]:\n    print(\"CORRECT PREDICTION !!!\")\nelse:\n    print(\"WRONG PREDICTION !!!\")","metadata":{"execution":{"iopub.status.busy":"2024-05-07T03:23:38.547762Z","iopub.execute_input":"2024-05-07T03:23:38.548060Z","iopub.status.idle":"2024-05-07T03:23:38.820568Z","shell.execute_reply.started":"2024-05-07T03:23:38.548034Z","shell.execute_reply":"2024-05-07T03:23:38.819643Z"},"trusted":true},"execution_count":null,"outputs":[]}]}