{"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":7445098,"sourceType":"datasetVersion","datasetId":4207557}],"dockerImageVersionId":30887,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:33.912435Z","iopub.execute_input":"2025-02-24T08:20:33.912755Z","iopub.status.idle":"2025-02-24T08:20:56.943658Z","shell.execute_reply.started":"2025-02-24T08:20:33.912733Z","shell.execute_reply":"2025-02-24T08:20:56.929141Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Enable auto-completion\n%config Completer.use_jedi = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:56.944872Z","iopub.execute_input":"2025-02-24T08:20:56.945196Z","iopub.status.idle":"2025-02-24T08:20:56.957828Z","shell.execute_reply.started":"2025-02-24T08:20:56.945158Z","shell.execute_reply":"2025-02-24T08:20:56.957090Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<div class=\"alert alert-info\" role=\"alert\">\n    In this 2nd project about <strong>\"The Cassava Leaf Disease\"</strong>, I would like to improve both my training and model performing by implementing:\n    <ul>\n        <li>Transfert Learning</li>\n        <li>Learning Rate Scheduling</li>\n        <li>Checkpointing</li>\n        <li>Early Stopping</li>\n    </ul>\n</div>","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport numpy as np\n\nimport torch\nimport PIL\nimport torchvision\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchinfo\n\nfrom torchvision import datasets, transforms, models\nfrom torch.utils.data import DataLoader, random_split\nfrom torch.optim.lr_scheduler import StepLR\nfrom torchinfo import summary\nfrom sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay\nfrom tqdm import tqdm\nfrom collections import Counter","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:56.959483Z","iopub.execute_input":"2025-02-24T08:20:56.959739Z","iopub.status.idle":"2025-02-24T08:20:56.972291Z","shell.execute_reply.started":"2025-02-24T08:20:56.959718Z","shell.execute_reply":"2025-02-24T08:20:56.971416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"torch version : \", torch.__version__)\nprint(\"torchvision version : \", torchvision.__version__)\nprint(\"torchinfo version : \", torchinfo.__version__)\nprint(\"numpy version : \", np.__version__)\nprint(\"matplotlib version : \", matplotlib.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:56.973720Z","iopub.execute_input":"2025-02-24T08:20:56.973963Z","iopub.status.idle":"2025-02-24T08:20:57.016137Z","shell.execute_reply.started":"2025-02-24T08:20:56.973944Z","shell.execute_reply":"2025-02-24T08:20:57.015541Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check if GPU is available, set device accordingly\nif torch.cuda.is_available():\n    device = \"cuda\"\nelif torch.backends.mps.is_available():\n    device = \"mps\"\nelse:\n    device = \"cpu\"\n\nprint(f\"We're using {device} device\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.017052Z","iopub.execute_input":"2025-02-24T08:20:57.017327Z","iopub.status.idle":"2025-02-24T08:20:57.031298Z","shell.execute_reply.started":"2025-02-24T08:20:57.017298Z","shell.execute_reply":"2025-02-24T08:20:57.030682Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Import and Read Data","metadata":{}},{"cell_type":"code","source":"# Create path to acces to the data\ndata_dir = os.path.join(\"/kaggle\", \"input\")\ntrain_dir = os.path.join(data_dir, \"data\")\n\nprint(train_dir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.032105Z","iopub.execute_input":"2025-02-24T08:20:57.032375Z","iopub.status.idle":"2025-02-24T08:20:57.047697Z","shell.execute_reply.started":"2025-02-24T08:20:57.032346Z","shell.execute_reply":"2025-02-24T08:20:57.046920Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create a list of classes\nclasses = os.listdir(train_dir)\nprint(\"classes length: \", len(classes))\nprint(classes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.048509Z","iopub.execute_input":"2025-02-24T08:20:57.048781Z","iopub.status.idle":"2025-02-24T08:20:57.063997Z","shell.execute_reply.started":"2025-02-24T08:20:57.048753Z","shell.execute_reply":"2025-02-24T08:20:57.063324Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Explore data","metadata":{}},{"cell_type":"code","source":"# Count Categories in each classe\ndef categories_counts(filepath, classes):\n\n    class_distribution_dict = dict()\n    \n    for classe in classes:\n        sub_dir = os.path.join(filepath, classe)\n        sub_ls = os.listdir(sub_dir)\n        len_sub = len(sub_ls)\n    \n        class_distribution_dict[classe] = len_sub\n    \n    # Convert class_distribution_dict to dataFrame\n    class_distribution = pd.DataFrame(\n        list(class_distribution_dict.items()),\n        columns = [\"Class\", \"Count\"]\n    ).set_index(\"Class\")\n    \n    return class_distribution","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.064645Z","iopub.execute_input":"2025-02-24T08:20:57.064824Z","iopub.status.idle":"2025-02-24T08:20:57.076193Z","shell.execute_reply.started":"2025-02-24T08:20:57.064808Z","shell.execute_reply":"2025-02-24T08:20:57.075604Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dist_class = categories_counts(train_dir, classes)\n\n# Visualize the distribution of classes\ndist_class.sort_values(\"Count\").plot(kind=\"bar\")\n\nplt.xlabel(\"Classes\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Class Distribution\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.078835Z","iopub.execute_input":"2025-02-24T08:20:57.079037Z","iopub.status.idle":"2025-02-24T08:20:57.275539Z","shell.execute_reply.started":"2025-02-24T08:20:57.079012Z","shell.execute_reply":"2025-02-24T08:20:57.274749Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Transform data","metadata":{}},{"cell_type":"code","source":"# Create a class to convert data mode in RGB\nclass ConvertToRGB():\n    def __call__(self, img):\n        if img.mode != \"RGB\":\n            img = img.convert(\"RGB\")\n        return img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.277254Z","iopub.execute_input":"2025-02-24T08:20:57.277504Z","iopub.status.idle":"2025-02-24T08:20:57.281074Z","shell.execute_reply.started":"2025-02-24T08:20:57.277482Z","shell.execute_reply":"2025-02-24T08:20:57.280329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create Transformer\ntransform = transforms.Compose(\n    [\n        ConvertToRGB(),\n        transforms.Resize((224, 224)),\n        transforms.ToTensor()\n    ]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.281816Z","iopub.execute_input":"2025-02-24T08:20:57.282005Z","iopub.status.idle":"2025-02-24T08:20:57.296049Z","shell.execute_reply.started":"2025-02-24T08:20:57.281988Z","shell.execute_reply":"2025-02-24T08:20:57.295214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Apply transformation to all data\ndataset = datasets.ImageFolder(root=train_dir, transform=transform)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.296740Z","iopub.execute_input":"2025-02-24T08:20:57.296985Z","iopub.status.idle":"2025-02-24T08:20:57.434551Z","shell.execute_reply.started":"2025-02-24T08:20:57.296961Z","shell.execute_reply":"2025-02-24T08:20:57.433926Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Batch data\nbatch_size = 32\ndata_loader = DataLoader(dataset, batch_size=batch_size)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.435278Z","iopub.execute_input":"2025-02-24T08:20:57.435513Z","iopub.status.idle":"2025-02-24T08:20:57.440376Z","shell.execute_reply.started":"2025-02-24T08:20:57.435493Z","shell.execute_reply":"2025-02-24T08:20:57.439671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# get the first batch\nfirst_batch = next(iter(data_loader))\nprint(\"Shape of batch: \", first_batch[0].shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.441111Z","iopub.execute_input":"2025-02-24T08:20:57.441383Z","iopub.status.idle":"2025-02-24T08:20:57.937472Z","shell.execute_reply.started":"2025-02-24T08:20:57.441363Z","shell.execute_reply":"2025-02-24T08:20:57.936683Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create a function to get mean and standard deviation\ndef get_mean_std(loader):\n    channel_sum, channel_sum_squared, num_batch = 0, 0, 0\n\n    for data, _ in tqdm(loader, desc=\"Computing mean and std\", leave=False):\n        channel_sum += torch.mean(data, dim=[0, 2, 3])\n        channel_sum_squared += torch.mean(data**2, dim=[0, 2, 3])\n        num_batch += 1\n\n        mean = channel_sum / num_batch\n        std = (channel_sum_squared / num_batch - mean**2) **0.5\n\n    return mean, std","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.938233Z","iopub.execute_input":"2025-02-24T08:20:57.938478Z","iopub.status.idle":"2025-02-24T08:20:57.943313Z","shell.execute_reply.started":"2025-02-24T08:20:57.938455Z","shell.execute_reply":"2025-02-24T08:20:57.942456Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mean, std = get_mean_std(data_loader)\nprint(mean, std)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:20:57.944126Z","iopub.execute_input":"2025-02-24T08:20:57.944329Z","iopub.status.idle":"2025-02-24T08:23:46.228968Z","shell.execute_reply.started":"2025-02-24T08:20:57.944310Z","shell.execute_reply":"2025-02-24T08:23:46.228109Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Normalize data\ntransform_norm = transforms.Compose(\n    [\n        ConvertToRGB(),\n        transforms.Resize((224, 224)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=mean, std=std)\n    ]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:23:46.234228Z","iopub.execute_input":"2025-02-24T08:23:46.234480Z","iopub.status.idle":"2025-02-24T08:23:46.238526Z","shell.execute_reply.started":"2025-02-24T08:23:46.234459Z","shell.execute_reply":"2025-02-24T08:23:46.237888Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Apply the new transformation to all data\nnorm_dataset = datasets.ImageFolder(root=train_dir, transform=transform_norm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:23:46.239294Z","iopub.execute_input":"2025-02-24T08:23:46.239523Z","iopub.status.idle":"2025-02-24T08:23:52.542141Z","shell.execute_reply.started":"2025-02-24T08:23:46.239504Z","shell.execute_reply":"2025-02-24T08:23:52.541474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split data\ng = torch.Generator()\ng.manual_seed(42)\n\ntrain_dataset, val_dataset = random_split(norm_dataset, [0.8, 0.2], generator=g)\n\npercent_train = np.round(100 * len(train_dataset) / len(norm_dataset), 2)\npercent_val = np.round(100 * len(val_dataset) / len(norm_dataset), 2)\n\nprint(f\"Train data is {percent_train}% of full data\")\nprint(f\"Validation data is {percent_val}% of full data\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:23:52.543108Z","iopub.execute_input":"2025-02-24T08:23:52.543413Z","iopub.status.idle":"2025-02-24T08:23:52.549766Z","shell.execute_reply.started":"2025-02-24T08:23:52.543368Z","shell.execute_reply":"2025-02-24T08:23:52.549004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create a function to count class\ndef class_counts(dataset):\n    c = Counter(x[1] for x in tqdm(dataset))\n    try:\n        class_to_index = dataset.class_to_idx\n    except AttributeError:\n        class_to_index = dataset.dataset.class_to_idx\n    return pd.Series({cat: c[idx] for cat, idx in class_to_index.items()})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:23:52.550410Z","iopub.execute_input":"2025-02-24T08:23:52.550670Z","iopub.status.idle":"2025-02-24T08:23:52.566528Z","shell.execute_reply.started":"2025-02-24T08:23:52.550649Z","shell.execute_reply":"2025-02-24T08:23:52.565474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_count= class_counts(train_dataset)\n\n# Visualize it\ntrain_count.plot(kind=\"bar\")\nplt.xlabel(\"Class\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Class Distribution: Training\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:23:52.567509Z","iopub.execute_input":"2025-02-24T08:23:52.567818Z","iopub.status.idle":"2025-02-24T08:26:28.842996Z","shell.execute_reply.started":"2025-02-24T08:23:52.567781Z","shell.execute_reply":"2025-02-24T08:26:28.842197Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_count = class_counts(val_dataset)\n\n# Visualize it\nval_count.plot(kind=\"bar\")\nplt.xlabel(\"Class\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Class Distribution: Validation\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:26:28.843812Z","iopub.execute_input":"2025-02-24T08:26:28.844103Z","iopub.status.idle":"2025-02-24T08:27:04.538307Z","shell.execute_reply.started":"2025-02-24T08:26:28.844071Z","shell.execute_reply":"2025-02-24T08:27:04.537560Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Batch train and val datasets\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, generator=g, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, generator=g, shuffle=False)\n\nprint(\"First batch shape: \", next(iter(train_loader))[0].shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:04.539181Z","iopub.execute_input":"2025-02-24T08:27:04.539499Z","iopub.status.idle":"2025-02-24T08:27:04.788812Z","shell.execute_reply.started":"2025-02-24T08:27:04.539475Z","shell.execute_reply":"2025-02-24T08:27:04.787966Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Build Model","metadata":{}},{"cell_type":"markdown","source":"#### *Choose the resnet50 model for Transfert Learning*\n","metadata":{}},{"cell_type":"code","source":"model = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\nprint(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:04.789672Z","iopub.execute_input":"2025-02-24T08:27:04.789932Z","iopub.status.idle":"2025-02-24T08:27:05.242354Z","shell.execute_reply.started":"2025-02-24T08:27:04.789902Z","shell.execute_reply":"2025-02-24T08:27:05.241467Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Freeze model so their parameters will not tweaked when training the model\nfor params in model.parameters():\n    params.requires_grad = False\n\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2025-02-24T08:27:05.243148Z","iopub.execute_input":"2025-02-24T08:27:05.243371Z","iopub.status.idle":"2025-02-24T08:27:05.249361Z","shell.execute_reply.started":"2025-02-24T08:27:05.243350Z","shell.execute_reply":"2025-02-24T08:27:05.248461Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get the name of the last model\nlist(model.named_modules())[-1]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.250242Z","iopub.execute_input":"2025-02-24T08:27:05.250532Z","iopub.status.idle":"2025-02-24T08:27:05.264332Z","shell.execute_reply.started":"2025-02-24T08:27:05.250505Z","shell.execute_reply":"2025-02-24T08:27:05.263672Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We'll build up a network with the following structure\n\n- Linear layer of 256 neurons\n- ReLU\n- Dropout\n- Linear layer of 5 neurons for output","metadata":{}},{"cell_type":"code","source":"in_features = model.fc.in_features\n\n# Create new model\nclassifier = nn.Sequential()\n\nclassifier.append(nn.Linear(in_features=in_features, out_features=256))\nclassifier.append(nn.ReLU())\nclassifier.append(nn.Dropout())\nclassifier.append(nn.Linear(256, 5))\n\nprint(classifier)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.267370Z","iopub.execute_input":"2025-02-24T08:27:05.267596Z","iopub.status.idle":"2025-02-24T08:27:05.281495Z","shell.execute_reply.started":"2025-02-24T08:27:05.267578Z","shell.execute_reply":"2025-02-24T08:27:05.280838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# No replace the model fc to classifier\nmodel.fc = classifier\n\nprint(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.283025Z","iopub.execute_input":"2025-02-24T08:27:05.283297Z","iopub.status.idle":"2025-02-24T08:27:05.291996Z","shell.execute_reply.started":"2025-02-24T08:27:05.283274Z","shell.execute_reply":"2025-02-24T08:27:05.291151Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training the model with Callbacks","metadata":{}},{"cell_type":"code","source":"# Compute the Loss Function\nloss_fn = nn.CrossEntropyLoss()\n\n# Calculate the optimizer\noptimizer = optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-4)\n\nmodel.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.292839Z","iopub.execute_input":"2025-02-24T08:27:05.293124Z","iopub.status.idle":"2025-02-24T08:27:05.350724Z","shell.execute_reply.started":"2025-02-24T08:27:05.293094Z","shell.execute_reply":"2025-02-24T08:27:05.350061Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### *Learning Rate Scheduling*\nwe'll use StepLR to create scheduler it decays the learning rate during training by multiplicative gamma","metadata":{}},{"cell_type":"code","source":"step_size = 5\ngamma = 0.1\n\nscheduler = StepLR(\n    optimizer, \n    step_size=step_size, \n    gamma=gamma,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.351423Z","iopub.execute_input":"2025-02-24T08:27:05.351612Z","iopub.status.idle":"2025-02-24T08:27:05.355079Z","shell.execute_reply.started":"2025-02-24T08:27:05.351596Z","shell.execute_reply":"2025-02-24T08:27:05.354431Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### *Early Stopping*","metadata":{}},{"cell_type":"code","source":"def early_stopping(validation_loss, best_val_loss, counter):\n    \"\"\"\n        The function checks:\n         - If the validation loss doesn't improve the counter will be 0, else we add 1\n         - If counter doesn't surpass 4 (counter=epochs)        \n    \"\"\"\n\n    stop = False\n\n    if validation_loss < best_val_loss:\n        counter = 0\n    else:\n        counter += 1\n\n    # Check if counter is >= patience (5 epochs in our case)\n    # Set stop variable accordingly\n    if counter >= 5:\n        stop = True\n\n    return counter, stop","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.355821Z","iopub.execute_input":"2025-02-24T08:27:05.356053Z","iopub.status.idle":"2025-02-24T08:27:05.368329Z","shell.execute_reply.started":"2025-02-24T08:27:05.356024Z","shell.execute_reply":"2025-02-24T08:27:05.367597Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### *Checkpointing*\nCheck if the validation loss improved. If yes, we will save the model","metadata":{}},{"cell_type":"code","source":"def checkpointing(validation_loss, best_val_loss, model, optimizer, save_path):\n\n    if validation_loss < best_val_loss:\n        torch.save(\n            {\n                \"model_state_dict\": model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"loss\": best_val_loss,\n            },\n            save_path,\n        )\n        print(f\"Checkpoint saved with validation loss {validation_loss:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.369070Z","iopub.execute_input":"2025-02-24T08:27:05.369314Z","iopub.status.idle":"2025-02-24T08:27:05.381716Z","shell.execute_reply.started":"2025-02-24T08:27:05.369295Z","shell.execute_reply":"2025-02-24T08:27:05.381142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(model, optimizer, loss_fn, data_loader, device=\"cpu\"):\n    training_loss = 0.0\n    model.train()\n\n    # Iterate over all batches in the training set to complete one epoch\n    for inputs, targets in tqdm(data_loader, desc=\"Training\", leave=False):\n        optimizer.zero_grad()\n        inputs = inputs.to(device)\n        targets = targets.to(device)\n\n        output = model(inputs)\n        loss = loss_fn(output, targets)\n\n        loss.backward()\n        optimizer.step()\n        training_loss += loss.data.item() * inputs.size(0)\n\n    return training_loss / len(data_loader.dataset)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.382302Z","iopub.execute_input":"2025-02-24T08:27:05.382506Z","iopub.status.idle":"2025-02-24T08:27:05.397245Z","shell.execute_reply.started":"2025-02-24T08:27:05.382489Z","shell.execute_reply":"2025-02-24T08:27:05.396635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def score(model, data_loader, loss_fn, device=\"cpu\"):\n    total_loss = 0\n    total_correct = 0\n\n    model.eval()\n    with torch.no_grad():\n        for inputs, targets in tqdm(data_loader, desc=\"Scoring\", leave=False):\n            inputs = inputs.to(device)\n            output = model(inputs)\n\n            targets = targets.to(device)\n            loss = loss_fn(output, targets)\n            total_loss += loss.data.item() * inputs.size(0)\n\n            correct = torch.eq(torch.argmax(output, dim=1), targets)\n            total_correct += torch.sum(correct).item()\n\n    n_observations = data_loader.batch_size * len(data_loader)\n    average_loss = total_loss / n_observations\n    accuracy = total_correct / n_observations\n    return average_loss, accuracy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.397940Z","iopub.execute_input":"2025-02-24T08:27:05.398141Z","iopub.status.idle":"2025-02-24T08:27:05.411012Z","shell.execute_reply.started":"2025-02-24T08:27:05.398122Z","shell.execute_reply":"2025-02-24T08:27:05.410237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(\n    model,\n    optimizer,\n    loss_fn,\n    train_loader,\n    val_loader,\n    epochs=20,\n    device=\"cpu\",\n    scheduler=None,\n    checkpoint_path=None,\n    early_stopping=None,\n):\n    # Track the model progress over epochs\n    train_losses = []\n    train_accuracies = []\n    val_losses = []\n    val_accuracies = []\n    learning_rates = []\n\n    # Create the trackers if needed for checkpointing and early stopping\n    best_val_loss = float(\"inf\")\n    early_stopping_counter = 0\n\n    print(\"Model evaluation before start of training...\")\n    # Test on training set\n    train_loss, train_accuracy = score(model, train_loader, loss_fn, device)\n    train_losses.append(train_loss)\n    train_accuracies.append(train_accuracy)\n    # Test on validation set\n    validation_loss, validation_accuracy = score(model, val_loader, loss_fn, device)\n    val_losses.append(validation_loss)\n    val_accuracies.append(validation_accuracy)\n\n    for epoch in range(1, epochs + 1):\n        print(\"\\n\")\n        print(f\"Starting epoch {epoch}/{epochs}\")\n\n        # Train one epoch\n        train_epoch(model, optimizer, loss_fn, train_loader, device)\n\n        # Evaluate training results\n        train_loss, train_accuracy = score(model, train_loader, loss_fn, device)\n        train_losses.append(train_loss)\n        train_accuracies.append(train_accuracy)\n\n        # Test on validation set\n        validation_loss, validation_accuracy = score(model, val_loader, loss_fn, device)\n        val_losses.append(validation_loss)\n        val_accuracies.append(validation_accuracy)\n\n        print(f\"Epoch: {epoch}\")\n        print(f\"Training loss: {train_loss:.4f}\")\n        print(f\"Training accuracy: {train_accuracy*100:.4f}%\")\n        print(f\"Validation loss: {validation_loss:.4f}\")\n        print(f\"Validation accuracy: {validation_accuracy*100:.4f}%\")\n\n        # # Log the learning rate and have the scheduler adjust it\n        lr = optimizer.param_groups[0][\"lr\"]\n        learning_rates.append(lr)\n        if scheduler:\n            scheduler.step()\n\n        # Checkpointing saves the model if current model is better than best so far\n        if checkpoint_path:\n            checkpointing(\n                validation_loss, best_val_loss, model, optimizer, checkpoint_path\n            )\n\n        # Early Stopping\n        if early_stopping:\n            early_stopping_counter, stop = early_stopping(\n                validation_loss, best_val_loss, early_stopping_counter\n            )\n            if stop:\n                print(f\"Early stopping triggered after {epoch} epochs\")\n                break\n\n        if validation_loss < best_val_loss:\n            best_val_loss = validation_loss\n\n    return (\n        learning_rates,\n        train_losses,\n        val_losses,\n        train_accuracies,\n        val_accuracies,\n        epoch,\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.411711Z","iopub.execute_input":"2025-02-24T08:27:05.411894Z","iopub.status.idle":"2025-02-24T08:27:05.430511Z","shell.execute_reply.started":"2025-02-24T08:27:05.411878Z","shell.execute_reply":"2025-02-24T08:27:05.429718Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoints_dir = os.path.join(\"/kaggle\", \"working\", \"model\")   \n\n# Define the checkpoint path within this writable directory  \ncheckpoint_path = os.path.join(checkpoints_dir, \"LR_model.pth\")\n\n# Create the 'checkpoints' directory if it doesn't exist  \nos.makedirs(os.path.dirname(checkpoint_path), exist_ok=True) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.431405Z","iopub.execute_input":"2025-02-24T08:27:05.431658Z","iopub.status.idle":"2025-02-24T08:27:05.448166Z","shell.execute_reply.started":"2025-02-24T08:27:05.431639Z","shell.execute_reply":"2025-02-24T08:27:05.447342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs_to_train = 50\n\ntrain_results = train(\n    model,\n    optimizer,\n    loss_fn,\n    train_loader,\n    val_loader,\n    epochs=epochs_to_train,\n    device=device,\n    scheduler=scheduler,\n    checkpoint_path=checkpoint_path,\n    early_stopping=early_stopping,\n)\n\n(\n    learning_rates,\n    train_losses,\n    valid_losses,\n    train_accuracies,\n    valid_accuracies,\n    epochs,\n) = train_results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T08:27:05.448996Z","iopub.execute_input":"2025-02-24T08:27:05.449263Z","iopub.status.idle":"2025-02-24T09:38:16.359189Z","shell.execute_reply.started":"2025-02-24T08:27:05.449236Z","shell.execute_reply":"2025-02-24T09:38:16.357867Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = torch.load(checkpoint_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T09:38:16.360159Z","iopub.execute_input":"2025-02-24T09:38:16.360482Z","iopub.status.idle":"2025-02-24T09:38:16.456245Z","shell.execute_reply.started":"2025-02-24T09:38:16.360449Z","shell.execute_reply":"2025-02-24T09:38:16.455328Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot train accuracies, use label=\"Training Accuracy\"\nplt.plot(train_accuracies, label=\"Training Accuracy\")\n# Plot validation accuracies, use label=\"Validation Accuracy\"\nplt.plot(valid_accuracies, label=\"Validation Accuracy\")\nplt.ylim([0, 1])\nplt.title(\"Accuracy over epochs\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Accuracy\")\nplt.legend();","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T09:38:16.457141Z","iopub.execute_input":"2025-02-24T09:38:16.457430Z","iopub.status.idle":"2025-02-24T09:38:16.651734Z","shell.execute_reply.started":"2025-02-24T09:38:16.457381Z","shell.execute_reply":"2025-02-24T09:38:16.651063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.plot(train_losses, label=\"Training Loss\")\nplt.plot(valid_losses, label=\"Validation Loss\")\nplt.ylim([0, 1.7])\nplt.title(\"Loss over epochs\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Loss\")\nplt.legend();","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T09:38:16.652413Z","iopub.execute_input":"2025-02-24T09:38:16.652692Z","iopub.status.idle":"2025-02-24T09:38:16.861376Z","shell.execute_reply.started":"2025-02-24T09:38:16.652671Z","shell.execute_reply":"2025-02-24T09:38:16.860584Z"}},"outputs":[],"execution_count":null}]}