SnehaAkula/case
0
1{2 "cells": [3 {4 "cell_type": "markdown",5 "metadata": {6 "id": "f__H59xsa0MS"7 },8 "source": [9 "### Load Dataset"10 ]11 },12 {13 "cell_type": "code",14 "execution_count": 1,15 "metadata": {16 "id": "BHi56mkNZs2h"17 },18 "outputs": [],19 "source": [20 "import pandas as pd"21 ]22 },23 {24 "cell_type": "code",25 "execution_count": null,26 "metadata": {27 "colab": {28 "base_uri": "https://localhost:8080/"29 },30 "id": "qduvy-i4yQCW",31 "outputId": "9f0fb200-84dd-465f-ab79-e9c91c6c3c86"32 },33 "outputs": [34 {35 "name": "stdout",36 "output_type": "stream",37 "text": [38 "Accuracy: 0.9567369876455109\n"39 ]40 }41 ],42 "source": [43 "from sklearn.feature_extraction.text import TfidfVectorizer\n",44 "from xgboost import XGBClassifier\n",45 "from sklearn.model_selection import train_test_split\n",46 "from sklearn.metrics import accuracy_score\n",47 "\n",48 "# Example dataset with 'input_sequence' and 'new_claim' columns\n",49 "X = full_data['input_sequence'] + \" \" + full_data['new_claim']\n",50 "full_data['target'] = full_data['target'].apply(lambda x: 1 if x == 'different_case' else 0)\n",51 "y = full_data['target'] # Labels (same_case, different_case)\n",52 "\n",53 "# Use TF-IDF to convert text into numerical features\n",54 "vectorizer = TfidfVectorizer(max_features=5000)\n",55 "X_transformed = vectorizer.fit_transform(X)\n",56 "\n",57 "# Split the data\n",58 "X_train, X_test, y_train, y_test = train_test_split(X_transformed, y, test_size=0.2, random_state=42)\n",59 "\n",60 "# Train the XGBoost model\n",61 "xgb_model = XGBClassifier()\n",62 "xgb_model.fit(X_train, y_train)\n",63 "\n",64 "# Predict and evaluate\n",65 "y_pred = xgb_model.predict(X_test)\n",66 "print(f\"Accuracy: {accuracy_score(y_test, y_pred)}\")"67 ]68 },69 {70 "cell_type": "code",71 "execution_count": null,72 "metadata": {73 "id": "Q13H-wNz3qBv"74 },75 "outputs": [],76 "source": [77 "bmark_df = pd.read_csv(\"/content/drive/MyDrive/auto_complete/data_v2/bmark_data.csv\")\n",78 "X_bmark = bmark_df['input_sequence'] + \" \" + bmark_df['new_claim']\n",79 "bmark_df['target'] = bmark_df['target'].apply(lambda x: 1 if x == 'different_case' else 0)\n",80 "y_bmark = bmark_df['target']"81 ]82 },83 {84 "cell_type": "code",85 "execution_count": null,86 "metadata": {87 "id": "Mj8_PS6x4LGI"88 },89 "outputs": [],90 "source": [91 "X_transformed = vectorizer.fit_transform(X_bmark)"92 ]93 },94 {95 "cell_type": "code",96 "execution_count": null,97 "metadata": {98 "colab": {99 "base_uri": "https://localhost:8080/"100 },101 "id": "L0tWuevU4RZ6",102 "outputId": "4ecde61b-0266-4118-b8e7-f7c0bacdd3a9"103 },104 "outputs": [105 {106 "name": "stdout",107 "output_type": "stream",108 "text": [109 "Accuracy: 0.509419983065199\n"110 ]111 }112 ],113 "source": [114 "y_pred_bmark = xgb_model.predict(X_transformed)\n",115 "print(f\"Accuracy: {accuracy_score(y_bmark, y_pred_bmark)}\")"116 ]117 },118 {119 "cell_type": "markdown",120 "metadata": {121 "id": "-zA881igbtIy"122 },123 "source": [124 "### Tokenize the dataset"125 ]126 },127 {128 "cell_type": "code",129 "execution_count": null,130 "metadata": {131 "id": "lP5U7kX2bHtE"132 },133 "outputs": [],134 "source": [135 "# from transformers import BertTokenizer\n",136 "# # from datasets import Dataset\n",137 "\n",138 "# # Initialize the tokenizer\n",139 "# tokenizer = BertTokenizer.from_pretrained('bert-base-uncased', use_fast=True)\n",140 "# # tokenizer = DistilBertTokenizer.from_pretrained('distilbert-base-uncased', use_fast=True)\n",141 "\n",142 "\n",143 "# # Define tokenization function\n",144 "# def tokenize_function(row):\n",145 "# return tokenizer(\n",146 "# row['input_sequence'],\n",147 "# row['new_claim'],\n",148 "# padding=\"max_length\",\n",149 "# truncation=True,\n",150 "# max_length=128\n",151 "# )\n",152 "\n",153 "# # Define chunk size (50,000 rows per chunk)\n",154 "# chunk_size = 20000\n",155 "\n",156 "# # Process data in chunks\n",157 "# for i in range(400000, len(full_data), chunk_size):\n",158 "# print(i)\n",159 "# chunk = full_data[i:i+chunk_size]\n",160 "\n",161 "# # Tokenize the chunk\n",162 "# tokenized_chunk = chunk.apply(lambda row: tokenize_function(row), axis=1).tolist()\n",163 "\n",164 "# # Convert to DataFrame and save as CSV\n",165 "# tokenized_df = pd.DataFrame(tokenized_chunk)\n",166 "# tokenized_df['target'] = chunk['target'].values\n",167 "# tokenized_df.to_csv(f'/content/drive/MyDrive/auto_complete/data_v2/tokenized_data_chunk_{i}.csv', index=False)\n",168 "\n",169 "# print(f\"Processed and saved chunk {i//chunk_size + 1}\")\n"170 ]171 },172 {173 "cell_type": "code",174 "source": [175 "from ast import literal_eval\n",176 "\n",177 "df1 = pd.read_csv('/content/drive/MyDrive/auto_complete/data_v2/tokenized_data_2.7lakh.csv', converters={'attention_mask': literal_eval,\n",178 " 'input_ids': literal_eval,\n",179 " 'token_type_ids': literal_eval})\n",180 "df2 = pd.read_csv('/content/drive/MyDrive/auto_complete/data_v2/tokenized_data_4lakh.csv', converters={'attention_mask': literal_eval,\n",181 " 'input_ids': literal_eval,\n",182 " 'token_type_ids': literal_eval})\n",183 "\n",184 "final_tokenized_data = pd.concat([df1, df2])\n",185 "final_tokenized_data"186 ],187 "metadata": {188 "id": "9hV7tYvIXzZ2",189 "colab": {190 "base_uri": "https://localhost:8080/",191 "height": 423192 },193 "outputId": "ded8b44e-576b-41dd-88b9-ea2b92b7aa97"194 },195 "execution_count": 16,196 "outputs": [197 {198 "output_type": "execute_result",199 "data": {200 "text/plain": [201 " attention_mask \\\n",202 "0 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",203 "1 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",204 "2 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",205 "3 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",206 "4 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",207 "... ... \n",208 "399995 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",209 "399996 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",210 "399997 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",211 "399998 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",212 "399999 [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ... \n",213 "\n",214 " input_ids \\\n",215 "0 [101, 11616, 2381, 1024, 9706, 2080, 1011, 189... \n",216 "1 [101, 11616, 2381, 1024, 1029, 21451, 3490, 73... \n",217 "2 [101, 11616, 2381, 1024, 6819, 22864, 11636, 1... \n",218 "3 [101, 11616, 2381, 1024, 24471, 2072, 1010, 43... \n",219 "4 [101, 11616, 2381, 1024, 10047, 23041, 3989, 1... \n",220 "... ... \n",221 "399995 [101, 11616, 2381, 1024, 1011, 16021, 5358, 62... \n",222 "399996 [101, 11616, 2381, 1024, 15255, 2132, 1001, 10... \n",223 "399997 [101, 11616, 2381, 1024, 2632, 17635, 7405, 10... \n",224 "399998 [101, 11616, 2381, 1024, 5472, 18153, 1011, 35... \n",225 "399999 [101, 11616, 2381, 1024, 3108, 17964, 1006, 19... \n",226 "\n",227 " token_type_ids target \n",228 "0 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... different_case \n",229 "1 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... different_case \n",230 "2 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... different_case \n",231 "3 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... different_case \n",232 "4 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... same_case \n",233 "... ... ... \n",234 "399995 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... different_case \n",235 "399996 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... different_case \n",236 "399997 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... same_case \n",237 "399998 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... same_case \n",238 "399999 [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ... same_case \n",239 "\n",240 "[672629 rows x 4 columns]"241 ],242 "text/html": [243 "\n",244 " <div id=\"df-c2225e23-bca1-4db5-857a-835f2fd4cb36\" class=\"colab-df-container\">\n",245 " <div>\n",246 "<style scoped>\n",247 " .dataframe tbody tr th:only-of-type {\n",248 " vertical-align: middle;\n",249 " }\n",250 "\n",251 " .dataframe tbody tr th {\n",252 " vertical-align: top;\n",253 " }\n",254 "\n",255 " .dataframe thead th {\n",256 " text-align: right;\n",257 " }\n",258 "</style>\n",259 "<table border=\"1\" class=\"dataframe\">\n",260 " <thead>\n",261 " <tr style=\"text-align: right;\">\n",262 " <th></th>\n",263 " <th>attention_mask</th>\n",264 " <th>input_ids</th>\n",265 " <th>token_type_ids</th>\n",266 " <th>target</th>\n",267 " </tr>\n",268 " </thead>\n",269 " <tbody>\n",270 " <tr>\n",271 " <th>0</th>\n",272 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",273 " <td>[101, 11616, 2381, 1024, 9706, 2080, 1011, 189...</td>\n",274 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",275 " <td>different_case</td>\n",276 " </tr>\n",277 " <tr>\n",278 " <th>1</th>\n",279 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",280 " <td>[101, 11616, 2381, 1024, 1029, 21451, 3490, 73...</td>\n",281 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",282 " <td>different_case</td>\n",283 " </tr>\n",284 " <tr>\n",285 " <th>2</th>\n",286 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",287 " <td>[101, 11616, 2381, 1024, 6819, 22864, 11636, 1...</td>\n",288 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",289 " <td>different_case</td>\n",290 " </tr>\n",291 " <tr>\n",292 " <th>3</th>\n",293 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",294 " <td>[101, 11616, 2381, 1024, 24471, 2072, 1010, 43...</td>\n",295 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",296 " <td>different_case</td>\n",297 " </tr>\n",298 " <tr>\n",299 " <th>4</th>\n",300 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",301 " <td>[101, 11616, 2381, 1024, 10047, 23041, 3989, 1...</td>\n",302 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",303 " <td>same_case</td>\n",304 " </tr>\n",305 " <tr>\n",306 " <th>...</th>\n",307 " <td>...</td>\n",308 " <td>...</td>\n",309 " <td>...</td>\n",310 " <td>...</td>\n",311 " </tr>\n",312 " <tr>\n",313 " <th>399995</th>\n",314 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",315 " <td>[101, 11616, 2381, 1024, 1011, 16021, 5358, 62...</td>\n",316 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",317 " <td>different_case</td>\n",318 " </tr>\n",319 " <tr>\n",320 " <th>399996</th>\n",321 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",322 " <td>[101, 11616, 2381, 1024, 15255, 2132, 1001, 10...</td>\n",323 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",324 " <td>different_case</td>\n",325 " </tr>\n",326 " <tr>\n",327 " <th>399997</th>\n",328 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",329 " <td>[101, 11616, 2381, 1024, 2632, 17635, 7405, 10...</td>\n",330 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",331 " <td>same_case</td>\n",332 " </tr>\n",333 " <tr>\n",334 " <th>399998</th>\n",335 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",336 " <td>[101, 11616, 2381, 1024, 5472, 18153, 1011, 35...</td>\n",337 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",338 " <td>same_case</td>\n",339 " </tr>\n",340 " <tr>\n",341 " <th>399999</th>\n",342 " <td>[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, ...</td>\n",343 " <td>[101, 11616, 2381, 1024, 3108, 17964, 1006, 19...</td>\n",344 " <td>[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...</td>\n",345 " <td>same_case</td>\n",346 " </tr>\n",347 " </tbody>\n",348 "</table>\n",349 "<p>672629 rows × 4 columns</p>\n",350 "</div>\n",351 " <div class=\"colab-df-buttons\">\n",352 "\n",353 " <div class=\"colab-df-container\">\n",354 " <button class=\"colab-df-convert\" onclick=\"convertToInteractive('df-c2225e23-bca1-4db5-857a-835f2fd4cb36')\"\n",355 " title=\"Convert this dataframe to an interactive table.\"\n",356 " style=\"display:none;\">\n",357 "\n",358 " <svg xmlns=\"http://www.w3.org/2000/svg\" height=\"24px\" viewBox=\"0 -960 960 960\">\n",359 " <path d=\"M120-120v-720h720v720H120Zm60-500h600v-160H180v160Zm220 220h160v-160H400v160Zm0 220h160v-160H400v160ZM180-400h160v-160H180v160Zm440 0h160v-160H620v160ZM180-180h160v-160H180v160Zm440 0h160v-160H620v160Z\"/>\n",360 " </svg>\n",361 " </button>\n",362 "\n",363 " <style>\n",364 " .colab-df-container {\n",365 " display:flex;\n",366 " gap: 12px;\n",367 " }\n",368 "\n",369 " .colab-df-convert {\n",370 " background-color: #E8F0FE;\n",371 " border: none;\n",372 " border-radius: 50%;\n",373 " cursor: pointer;\n",374 " display: none;\n",375 " fill: #1967D2;\n",376 " height: 32px;\n",377 " padding: 0 0 0 0;\n",378 " width: 32px;\n",379 " }\n",380 "\n",381 " .colab-df-convert:hover {\n",382 " background-color: #E2EBFA;\n",383 " box-shadow: 0px 1px 2px rgba(60, 64, 67, 0.3), 0px 1px 3px 1px rgba(60, 64, 67, 0.15);\n",384 " fill: #174EA6;\n",385 " }\n",386 "\n",387 " .colab-df-buttons div {\n",388 " margin-bottom: 4px;\n",389 " }\n",390 "\n",391 " [theme=dark] .colab-df-convert {\n",392 " background-color: #3B4455;\n",393 " fill: #D2E3FC;\n",394 " }\n",395 "\n",396 " [theme=dark] .colab-df-convert:hover {\n",397 " background-color: #434B5C;\n",398 " box-shadow: 0px 1px 3px 1px rgba(0, 0, 0, 0.15);\n",399 " filter: drop-shadow(0px 1px 2px rgba(0, 0, 0, 0.3));\n",400 " fill: #FFFFFF;\n",401 " }\n",402 " </style>\n",403 "\n",404 " <script>\n",405 " const buttonEl =\n",406 " document.querySelector('#df-c2225e23-bca1-4db5-857a-835f2fd4cb36 button.colab-df-convert');\n",407 " buttonEl.style.display =\n",408 " google.colab.kernel.accessAllowed ? 'block' : 'none';\n",409 "\n",410 " async function convertToInteractive(key) {\n",411 " const element = document.querySelector('#df-c2225e23-bca1-4db5-857a-835f2fd4cb36');\n",412 " const dataTable =\n",413 " await google.colab.kernel.invokeFunction('convertToInteractive',\n",414 " [key], {});\n",415 " if (!dataTable) return;\n",416 "\n",417 " const docLinkHtml = 'Like what you see? Visit the ' +\n",418 " '<a target=\"_blank\" href=https://colab.research.google.com/notebooks/data_table.ipynb>data table notebook</a>'\n",419 " + ' to learn more about interactive tables.';\n",420 " element.innerHTML = '';\n",421 " dataTable['output_type'] = 'display_data';\n",422 " await google.colab.output.renderOutput(dataTable, element);\n",423 " const docLink = document.createElement('div');\n",424 " docLink.innerHTML = docLinkHtml;\n",425 " element.appendChild(docLink);\n",426 " }\n",427 " </script>\n",428 " </div>\n",429 "\n",430 "\n",431 "<div id=\"df-a07882c4-8037-4ca8-b537-79c0834fea02\">\n",432 " <button class=\"colab-df-quickchart\" onclick=\"quickchart('df-a07882c4-8037-4ca8-b537-79c0834fea02')\"\n",433 " title=\"Suggest charts\"\n",434 " style=\"display:none;\">\n",435 "\n",436 "<svg xmlns=\"http://www.w3.org/2000/svg\" height=\"24px\"viewBox=\"0 0 24 24\"\n",437 " width=\"24px\">\n",438 " <g>\n",439 " <path d=\"M19 3H5c-1.1 0-2 .9-2 2v14c0 1.1.9 2 2 2h14c1.1 0 2-.9 2-2V5c0-1.1-.9-2-2-2zM9 17H7v-7h2v7zm4 0h-2V7h2v10zm4 0h-2v-4h2v4z\"/>\n",440 " </g>\n",441 "</svg>\n",442 " </button>\n",443 "\n",444 "<style>\n",445 " .colab-df-quickchart {\n",446 " --bg-color: #E8F0FE;\n",447 " --fill-color: #1967D2;\n",448 " --hover-bg-color: #E2EBFA;\n",449 " --hover-fill-color: #174EA6;\n",450 " --disabled-fill-color: #AAA;\n",451 " --disabled-bg-color: #DDD;\n",452 " }\n",453 "\n",454 " [theme=dark] .colab-df-quickchart {\n",455 " --bg-color: #3B4455;\n",456 " --fill-color: #D2E3FC;\n",457 " --hover-bg-color: #434B5C;\n",458 " --hover-fill-color: #FFFFFF;\n",459 " --disabled-bg-color: #3B4455;\n",460 " --disabled-fill-color: #666;\n",461 " }\n",462 "\n",463 " .colab-df-quickchart {\n",464 " background-color: var(--bg-color);\n",465 " border: none;\n",466 " border-radius: 50%;\n",467 " cursor: pointer;\n",468 " display: none;\n",469 " fill: var(--fill-color);\n",470 " height: 32px;\n",471 " padding: 0;\n",472 " width: 32px;\n",473 " }\n",474 "\n",475 " .colab-df-quickchart:hover {\n",476 " background-color: var(--hover-bg-color);\n",477 " box-shadow: 0 1px 2px rgba(60, 64, 67, 0.3), 0 1px 3px 1px rgba(60, 64, 67, 0.15);\n",478 " fill: var(--button-hover-fill-color);\n",479 " }\n",480 "\n",481 " .colab-df-quickchart-complete:disabled,\n",482 " .colab-df-quickchart-complete:disabled:hover {\n",483 " background-color: var(--disabled-bg-color);\n",484 " fill: var(--disabled-fill-color);\n",485 " box-shadow: none;\n",486 " }\n",487 "\n",488 " .colab-df-spinner {\n",489 " border: 2px solid var(--fill-color);\n",490 " border-color: transparent;\n",491 " border-bottom-color: var(--fill-color);\n",492 " animation:\n",493 " spin 1s steps(1) infinite;\n",494 " }\n",495 "\n",496 " @keyframes spin {\n",497 " 0% {\n",498 " border-color: transparent;\n",499 " border-bottom-color: var(--fill-color);\n",500 " border-left-color: var(--fill-color);\n",501 " }\n",502 " 20% {\n",503 " border-color: transparent;\n",504 " border-left-color: var(--fill-color);\n",505 " border-top-color: var(--fill-color);\n",506 " }\n",507 " 30% {\n",508 " border-color: transparent;\n",509 " border-left-color: var(--fill-color);\n",510 " border-top-color: var(--fill-color);\n",511 " border-right-color: var(--fill-color);\n",512 " }\n",513 " 40% {\n",514 " border-color: transparent;\n",515 " border-right-color: var(--fill-color);\n",516 " border-top-color: var(--fill-color);\n",517 " }\n",518 " 60% {\n",519 " border-color: transparent;\n",520 " border-right-color: var(--fill-color);\n",521 " }\n",522 " 80% {\n",523 " border-color: transparent;\n",524 " border-right-color: var(--fill-color);\n",525 " border-bottom-color: var(--fill-color);\n",526 " }\n",527 " 90% {\n",528 " border-color: transparent;\n",529 " border-bottom-color: var(--fill-color);\n",530 " }\n",531 " }\n",532 "</style>\n",533 "\n",534 " <script>\n",535 " async function quickchart(key) {\n",536 " const quickchartButtonEl =\n",537 " document.querySelector('#' + key + ' button');\n",538 " quickchartButtonEl.disabled = true; // To prevent multiple clicks.\n",539 " quickchartButtonEl.classList.add('colab-df-spinner');\n",540 " try {\n",541 " const charts = await google.colab.kernel.invokeFunction(\n",542 " 'suggestCharts', [key], {});\n",543 " } catch (error) {\n",544 " console.error('Error during call to suggestCharts:', error);\n",545 " }\n",546 " quickchartButtonEl.classList.remove('colab-df-spinner');\n",547 " quickchartButtonEl.classList.add('colab-df-quickchart-complete');\n",548 " }\n",549 " (() => {\n",550 " let quickchartButtonEl =\n",551 " document.querySelector('#df-a07882c4-8037-4ca8-b537-79c0834fea02 button');\n",552 " quickchartButtonEl.style.display =\n",553 " google.colab.kernel.accessAllowed ? 'block' : 'none';\n",554 " })();\n",555 " </script>\n",556 "</div>\n",557 "\n",558 " <div id=\"id_edaa86bd-59bd-4f4b-8e4a-c22a362855b2\">\n",559 " <style>\n",560 " .colab-df-generate {\n",561 " background-color: #E8F0FE;\n",562 " border: none;\n",563 " border-radius: 50%;\n",564 " cursor: pointer;\n",565 " display: none;\n",566 " fill: #1967D2;\n",567 " height: 32px;\n",568 " padding: 0 0 0 0;\n",569 " width: 32px;\n",570 " }\n",571 "\n",572 " .colab-df-generate:hover {\n",573 " background-color: #E2EBFA;\n",574 " box-shadow: 0px 1px 2px rgba(60, 64, 67, 0.3), 0px 1px 3px 1px rgba(60, 64, 67, 0.15);\n",575 " fill: #174EA6;\n",576 " }\n",577 "\n",578 " [theme=dark] .colab-df-generate {\n",579 " background-color: #3B4455;\n",580 " fill: #D2E3FC;\n",581 " }\n",582 "\n",583 " [theme=dark] .colab-df-generate:hover {\n",584 " background-color: #434B5C;\n",585 " box-shadow: 0px 1px 3px 1px rgba(0, 0, 0, 0.15);\n",586 " filter: drop-shadow(0px 1px 2px rgba(0, 0, 0, 0.3));\n",587 " fill: #FFFFFF;\n",588 " }\n",589 " </style>\n",590 " <button class=\"colab-df-generate\" onclick=\"generateWithVariable('final_tokenized_data')\"\n",591 " title=\"Generate code using this dataframe.\"\n",592 " style=\"display:none;\">\n",593 "\n",594 " <svg xmlns=\"http://www.w3.org/2000/svg\" height=\"24px\"viewBox=\"0 0 24 24\"\n",595 " width=\"24px\">\n",596 " <path d=\"M7,19H8.4L18.45,9,17,7.55,7,17.6ZM5,21V16.75L18.45,3.32a2,2,0,0,1,2.83,0l1.4,1.43a1.91,1.91,0,0,1,.58,1.4,1.91,1.91,0,0,1-.58,1.4L9.25,21ZM18.45,9,17,7.55Zm-12,3A5.31,5.31,0,0,0,4.9,8.1,5.31,5.31,0,0,0,1,6.5,5.31,5.31,0,0,0,4.9,4.9,5.31,5.31,0,0,0,6.5,1,5.31,5.31,0,0,0,8.1,4.9,5.31,5.31,0,0,0,12,6.5,5.46,5.46,0,0,0,6.5,12Z\"/>\n",597 " </svg>\n",598 " </button>\n",599 " <script>\n",600 " (() => {\n",601 " const buttonEl =\n",602 " document.querySelector('#id_edaa86bd-59bd-4f4b-8e4a-c22a362855b2 button.colab-df-generate');\n",603 " buttonEl.style.display =\n",604 " google.colab.kernel.accessAllowed ? 'block' : 'none';\n",605 "\n",606 " buttonEl.onclick = () => {\n",607 " google.colab.notebook.generateWithVariable('final_tokenized_data');\n",608 " }\n",609 " })();\n",610 " </script>\n",611 " </div>\n",612 "\n",613 " </div>\n",614 " </div>\n"615 ],616 "application/vnd.google.colaboratory.intrinsic+json": {617 "type": "dataframe",618 "variable_name": "final_tokenized_data"619 }620 },621 "metadata": {},622 "execution_count": 16623 }624 ]625 },626 {627 "cell_type": "code",628 "source": [629 "final_tokenized_data.rename(columns={'target': 'labels'}, inplace=True)"630 ],631 "metadata": {632 "id": "xb5EXQEwozjp"633 },634 "execution_count": 20,635 "outputs": []636 },637 {638 "cell_type": "code",639 "source": [640 "from sklearn.model_selection import train_test_split\n",641 "\n",642 "# Convert the target column to numeric labels\n",643 "final_tokenized_data['labels'] = final_tokenized_data['labels'].apply(lambda x: int(1) if x == 'different_case' else int(0))\n",644 "\n",645 "# Split the data into training and validation sets\n",646 "train_data, val_data = train_test_split(final_tokenized_data, test_size=0.2, random_state=1)"647 ],648 "metadata": {649 "id": "mhvwBJyAmRIJ"650 },651 "execution_count": 26,652 "outputs": []653 },654 {655 "cell_type": "code",656 "source": [657 "import ast\n",658 "import numpy as np\n",659 "\n",660 "# Define a function to apply to each cell\n",661 "def convert_to_list(x):\n",662 " if isinstance(x, str):\n",663 " # Check if the string represents 'nan'\n",664 " if x.lower() == 'nan':\n",665 " return np.nan # Return np.nan to represent missing values\n",666 " try:\n",667 " # Convert string representation of a list to an actual list\n",668 " return ast.literal_eval(x)\n",669 " except (ValueError, SyntaxError):\n",670 " return x # If conversion fails, return the original value\n",671 " return x # If it's not a string, return the value as is\n",672 "\n",673 "# Apply the conversion function to each relevant column\n",674 "columns_to_convert = ['attention_mask', 'input_ids', 'token_type_ids']\n",675 "\n",676 "for col in columns_to_convert:\n",677 " train_data[col] = train_data[col].apply(convert_to_list)\n",678 "\n",679 "for col in columns_to_convert:\n",680 " val_data[col] = val_data[col].apply(convert_to_list)"681 ],682 "metadata": {683 "id": "R31OADhbqax7"684 },685 "execution_count": 36,686 "outputs": []687 },688 {689 "cell_type": "code",690 "source": [691 "from torch.utils.data import Dataset\n",692 "import torch\n",693 "\n",694 "class CustomDataset(Dataset):\n",695 " def __init__(self, encodings, labels):\n",696 " self.encodings = encodings\n",697 " self.labels = labels\n",698 "\n",699 " def __getitem__(self, idx):\n",700 " item = {key: torch.tensor(val[idx]) for key, val in self.encodings.items()}\n",701 " item['labels'] = torch.tensor(self.labels[idx])\n",702 " return item\n",703 "\n",704 " def __len__(self):\n",705 " return len(self.labels)\n",706 "\n",707 "# Convert tokenized_data to the necessary format\n",708 "def convert_to_dataset(df):\n",709 " encodings = {\n",710 " 'input_ids': df['input_ids'].tolist(),\n",711 " 'attention_mask': df['attention_mask'].tolist(),\n",712 " }\n",713 " labels = df['labels'].tolist()\n",714 " return CustomDataset(encodings, labels)\n",715 "\n",716 "# Prepare datasets\n",717 "train_dataset = convert_to_dataset(train_data)\n",718 "val_dataset = convert_to_dataset(val_data)\n"719 ],720 "metadata": {721 "id": "xp5l6NSDo-Ud"722 },723 "execution_count": 37,724 "outputs": []725 },726 {727 "cell_type": "code",728 "source": [729 "from transformers import BertForSequenceClassification, Trainer, TrainingArguments\n",730 "\n",731 "# Initialize the model\n",732 "model = BertForSequenceClassification.from_pretrained('bert-base-uncased', num_labels=2)\n",733 "\n",734 "# Set up training arguments\n",735 "training_args = TrainingArguments(\n",736 " output_dir='./results',\n",737 " num_train_epochs=3,\n",738 " per_device_train_batch_size=8,\n",739 " per_device_eval_batch_size=8,\n",740 " warmup_steps=500,\n",741 " weight_decay=0.01,\n",742 " logging_dir='./logs',\n",743 " logging_steps=10,\n",744 ")\n",745 "\n",746 "# Set up the Trainer\n",747 "trainer = Trainer(\n",748 " model=model,\n",749 " args=training_args,\n",750 " train_dataset=train_dataset,\n",751 " eval_dataset=val_dataset,\n",752 ")\n",753 "\n",754 "# Train the model\n",755 "trainer.train()"756 ],757 "metadata": {758 "colab": {759 "base_uri": "https://localhost:8080/",760 "height": 1000761 },762 "id": "FrllxDx4o_gG",763 "outputId": "bb363c5d-9c57-460f-9a49-f5a1a221eac6"764 },765 "execution_count": null,766 "outputs": [767 {768 "metadata": {769 "tags": null770 },771 "name": "stderr",772 "output_type": "stream",773 "text": [774 "Some weights of BertForSequenceClassification were not initialized from the model checkpoint at bert-base-uncased and are newly initialized: ['classifier.bias', 'classifier.weight']\n",775 "You should probably TRAIN this model on a down-stream task to be able to use it for predictions and inference.\n"776 ]777 },778 {779 "data": {780 "text/html": [781 "\n",782 " <div>\n",783 " \n",784 " <progress value='7813' max='201789' style='width:300px; height:20px; vertical-align: middle;'></progress>\n",785 " [ 7813/201789 37:28 < 15:30:47, 3.47 it/s, Epoch 0.12/3]\n",786 " </div>\n",787 " <table border=\"1\" class=\"dataframe\">\n",788 " <thead>\n",789 " <tr style=\"text-align: left;\">\n",790 " <th>Step</th>\n",791 " <th>Training Loss</th>\n",792 " </tr>\n",793 " </thead>\n",794 " <tbody>\n",795 " <tr>\n",796 " <td>10</td>\n",797 " <td>0.677500</td>\n",798 " </tr>\n",799 " <tr>\n",800 " <td>20</td>\n",801 " <td>0.690600</td>\n",802 " </tr>\n",803 " <tr>\n",804 " <td>30</td>\n",805 " <td>0.648200</td>\n",806 " </tr>\n",807 " <tr>\n",808 " <td>40</td>\n",809 " <td>0.607200</td>\n",810 " </tr>\n",811 " <tr>\n",812 " <td>50</td>\n",813 " <td>0.573300</td>\n",814 " </tr>\n",815 " <tr>\n",816 " <td>60</td>\n",817 " <td>0.508200</td>\n",818 " </tr>\n",819 " <tr>\n",820 " <td>70</td>\n",821 " <td>0.434200</td>\n",822 " </tr>\n",823 " <tr>\n",824 " <td>80</td>\n",825 " <td>0.402900</td>\n",826 " </tr>\n",827 " <tr>\n",828 " <td>90</td>\n",829 " <td>0.455000</td>\n",830 " </tr>\n",831 " <tr>\n",832 " <td>100</td>\n",833 " <td>0.363800</td>\n",834 " </tr>\n",835 " <tr>\n",836 " <td>110</td>\n",837 " <td>0.342000</td>\n",838 " </tr>\n",839 " <tr>\n",840 " <td>120</td>\n",841 " <td>0.246800</td>\n",842 " </tr>\n",843 " <tr>\n",844 " <td>130</td>\n",845 " <td>0.246900</td>\n",846 " </tr>\n",847 " <tr>\n",848 " <td>140</td>\n",849 " <td>0.309200</td>\n",850 " </tr>\n",851 " <tr>\n",852 " <td>150</td>\n",853 " <td>0.317600</td>\n",854 " </tr>\n",855 " <tr>\n",856 " <td>160</td>\n",857 " <td>0.205000</td>\n",858 " </tr>\n",859 " <tr>\n",860 " <td>170</td>\n",861 " <td>0.224900</td>\n",862 " </tr>\n",863 " <tr>\n",864 " <td>180</td>\n",865 " <td>0.222100</td>\n",866 " </tr>\n",867 " <tr>\n",868 " <td>190</td>\n",869 " <td>0.289100</td>\n",870 " </tr>\n",871 " <tr>\n",872 " <td>200</td>\n",873 " <td>0.350800</td>\n",874 " </tr>\n",875 " <tr>\n",876 " <td>210</td>\n",877 " <td>0.275000</td>\n",878 " </tr>\n",879 " <tr>\n",880 " <td>220</td>\n",881 " <td>0.320900</td>\n",882 " </tr>\n",883 " <tr>\n",884 " <td>230</td>\n",885 " <td>0.189500</td>\n",886 " </tr>\n",887 " <tr>\n",888 " <td>240</td>\n",889 " <td>0.310700</td>\n",890 " </tr>\n",891 " <tr>\n",892 " <td>250</td>\n",893 " <td>0.267200</td>\n",894 " </tr>\n",895 " <tr>\n",896 " <td>260</td>\n",897 " <td>0.148200</td>\n",898 " </tr>\n",899 " <tr>\n",900 " <td>270</td>\n",901 " <td>0.187900</td>\n",902 " </tr>\n",903 " <tr>\n",904 " <td>280</td>\n",905 " <td>0.203800</td>\n",906 " </tr>\n",907 " <tr>\n",908 " <td>290</td>\n",909 " <td>0.257600</td>\n",910 " </tr>\n",911 " <tr>\n",912 " <td>300</td>\n",913 " <td>0.208600</td>\n",914 " </tr>\n",915 " <tr>\n",916 " <td>310</td>\n",917 " <td>0.292700</td>\n",918 " </tr>\n",919 " <tr>\n",920 " <td>320</td>\n",921 " <td>0.274900</td>\n",922 " </tr>\n",923 " <tr>\n",924 " <td>330</td>\n",925 " <td>0.344200</td>\n",926 " </tr>\n",927 " <tr>\n",928 " <td>340</td>\n",929 " <td>0.279700</td>\n",930 " </tr>\n",931 " <tr>\n",932 " <td>350</td>\n",933 " <td>0.198800</td>\n",934 " </tr>\n",935 " <tr>\n",936 " <td>360</td>\n",937 " <td>0.181900</td>\n",938 " </tr>\n",939 " <tr>\n",940 " <td>370</td>\n",941 " <td>0.300700</td>\n",942 " </tr>\n",943 " <tr>\n",944 " <td>380</td>\n",945 " <td>0.428800</td>\n",946 " </tr>\n",947 " <tr>\n",948 " <td>390</td>\n",949 " <td>0.198700</td>\n",950 " </tr>\n",951 " <tr>\n",952 " <td>400</td>\n",953 " <td>0.303900</td>\n",954 " </tr>\n",955 " <tr>\n",956 " <td>410</td>\n",957 " <td>0.194000</td>\n",958 " </tr>\n",959 " <tr>\n",960 " <td>420</td>\n",961 " <td>0.310000</td>\n",962 " </tr>\n",963 " <tr>\n",964 " <td>430</td>\n",965 " <td>0.256700</td>\n",966 " </tr>\n",967 " <tr>\n",968 " <td>440</td>\n",969 " <td>0.264200</td>\n",970 " </tr>\n",971 " <tr>\n",972 " <td>450</td>\n",973 " <td>0.250300</td>\n",974 " </tr>\n",975 " <tr>\n",976 " <td>460</td>\n",977 " <td>0.101900</td>\n",978 " </tr>\n",979 " <tr>\n",980 " <td>470</td>\n",981 " <td>0.147400</td>\n",982 " </tr>\n",983 " <tr>\n",984 " <td>480</td>\n",985 " <td>0.104800</td>\n",986 " </tr>\n",987 " <tr>\n",988 " <td>490</td>\n",989 " <td>0.189100</td>\n",990 " </tr>\n",991 " <tr>\n",992 " <td>500</td>\n",993 " <td>0.263800</td>\n",994 " </tr>\n",995 " <tr>\n",996 " <td>510</td>\n",997 " <td>0.176700</td>\n",998 " </tr>\n",999 " <tr>\n",1000 " <td>520</td>\n",1001 " <td>0.375200</td>\n",1002 " </tr>\n",1003 " <tr>\n",1004 " <td>530</td>\n",1005 " <td>0.342300</td>\n",1006 " </tr>\n",1007 " <tr>\n",1008 " <td>540</td>\n",1009 " <td>0.333000</td>\n",1010 " </tr>\n",1011 " <tr>\n",1012 " <td>550</td>\n",1013 " <td>0.288800</td>\n",1014 " </tr>\n",1015 " <tr>\n",1016 " <td>560</td>\n",1017 " <td>0.182300</td>\n",1018 " </tr>\n",1019 " <tr>\n",1020 " <td>570</td>\n",1021 " <td>0.296100</td>\n",1022 " </tr>\n",1023 " <tr>\n",1024 " <td>580</td>\n",1025 " <td>0.135800</td>\n",1026 " </tr>\n",1027 " <tr>\n",1028 " <td>590</td>\n",1029 " <td>0.168600</td>\n",1030 " </tr>\n",1031 " <tr>\n",1032 " <td>600</td>\n",1033 " <td>0.399600</td>\n",1034 " </tr>\n",1035 " <tr>\n",1036 " <td>610</td>\n",1037 " <td>0.476400</td>\n",1038 " </tr>\n",1039 " <tr>\n",1040 " <td>620</td>\n",1041 " <td>0.099200</td>\n",1042 " </tr>\n",1043 " <tr>\n",1044 " <td>630</td>\n",1045 " <td>0.295800</td>\n",1046 " </tr>\n",1047 " <tr>\n",1048 " <td>640</td>\n",1049 " <td>0.205300</td>\n",1050 " </tr>\n",1051 " <tr>\n",1052 " <td>650</td>\n",1053 " <td>0.104500</td>\n",1054 " </tr>\n",1055 " <tr>\n",1056 " <td>660</td>\n",1057 " <td>0.518100</td>\n",1058 " </tr>\n",1059 " <tr>\n",1060 " <td>670</td>\n",1061 " <td>0.413000</td>\n",1062 " </tr>\n",1063 " <tr>\n",1064 " <td>680</td>\n",1065 " <td>0.200900</td>\n",1066 " </tr>\n",1067 " <tr>\n",1068 " <td>690</td>\n",1069 " <td>0.184200</td>\n",1070 " </tr>\n",1071 " <tr>\n",1072 " <td>700</td>\n",1073 " <td>0.311800</td>\n",1074 " </tr>\n",1075 " <tr>\n",1076 " <td>710</td>\n",1077 " <td>0.281500</td>\n",1078 " </tr>\n",1079 " <tr>\n",1080 " <td>720</td>\n",1081 " <td>0.211200</td>\n",1082 " </tr>\n",1083 " <tr>\n",1084 " <td>730</td>\n",1085 " <td>0.283900</td>\n",1086 " </tr>\n",1087 " <tr>\n",1088 " <td>740</td>\n",1089 " <td>0.163700</td>\n",1090 " </tr>\n",1091 " <tr>\n",1092 " <td>750</td>\n",1093 " <td>0.340400</td>\n",1094 " </tr>\n",1095 " <tr>\n",1096 " <td>760</td>\n",1097 " <td>0.122900</td>\n",1098 " </tr>\n",1099 " <tr>\n",1100 " <td>770</td>\n",1101 " <td>0.154600</td>\n",1102 " </tr>\n",1103 " <tr>\n",1104 " <td>780</td>\n",1105 " <td>0.291900</td>\n",1106 " </tr>\n",1107 " <tr>\n",1108 " <td>790</td>\n",1109 " <td>0.296900</td>\n",1110 " </tr>\n",1111 " <tr>\n",1112 " <td>800</td>\n",1113 " <td>0.086700</td>\n",1114 " </tr>\n",1115 " <tr>\n",1116 " <td>810</td>\n",1117 " <td>0.185600</td>\n",1118 " </tr>\n",1119 " <tr>\n",1120 " <td>820</td>\n",1121 " <td>0.356500</td>\n",1122 " </tr>\n",1123 " <tr>\n",1124 " <td>830</td>\n",1125 " <td>0.294200</td>\n",1126 " </tr>\n",1127 " <tr>\n",1128 " <td>840</td>\n",1129 " <td>0.388700</td>\n",1130 " </tr>\n",1131 " <tr>\n",1132 " <td>850</td>\n",1133 " <td>0.415400</td>\n",1134 " </tr>\n",1135 " <tr>\n",1136 " <td>860</td>\n",1137 " <td>0.397500</td>\n",1138 " </tr>\n",1139 " <tr>\n",1140 " <td>870</td>\n",1141 " <td>0.188500</td>\n",1142 " </tr>\n",1143 " <tr>\n",1144 " <td>880</td>\n",1145 " <td>0.265600</td>\n",1146 " </tr>\n",1147 " <tr>\n",1148 " <td>890</td>\n",1149 " <td>0.305200</td>\n",1150 " </tr>\n",1151 " <tr>\n",1152 " <td>900</td>\n",1153 " <td>0.184000</td>\n",1154 " </tr>\n",1155 " <tr>\n",1156 " <td>910</td>\n",1157 " <td>0.178200</td>\n",1158 " </tr>\n",1159 " <tr>\n",1160 " <td>920</td>\n",1161 " <td>0.206100</td>\n",1162 " </tr>\n",1163 " <tr>\n",1164 " <td>930</td>\n",1165 " <td>0.102100</td>\n",1166 " </tr>\n",1167 " <tr>\n",1168 " <td>940</td>\n",1169 " <td>0.297800</td>\n",1170 " </tr>\n",1171 " <tr>\n",1172 " <td>950</td>\n",1173 " <td>0.270100</td>\n",1174 " </tr>\n",1175 " <tr>\n",1176 " <td>960</td>\n",1177 " <td>0.251800</td>\n",1178 " </tr>\n",1179 " <tr>\n",1180 " <td>970</td>\n",1181 " <td>0.252500</td>\n",1182 " </tr>\n",1183 " <tr>\n",1184 " <td>980</td>\n",1185 " <td>0.179700</td>\n",1186 " </tr>\n",1187 " <tr>\n",1188 " <td>990</td>\n",1189 " <td>0.249700</td>\n",1190 " </tr>\n",1191 " <tr>\n",1192 " <td>1000</td>\n",1193 " <td>0.117000</td>\n",1194 " </tr>\n",1195 " <tr>\n",1196 " <td>1010</td>\n",1197 " <td>0.340700</td>\n",1198 " </tr>\n",1199 " <tr>\n",1200 " <td>1020</td>\n",