{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchvision import datasets, transforms, models","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-03-22T18:31:28.418093Z","iopub.execute_input":"2023-03-22T18:31:28.418505Z","iopub.status.idle":"2023-03-22T18:31:28.42417Z","shell.execute_reply.started":"2023-03-22T18:31:28.41846Z","shell.execute_reply":"2023-03-22T18:31:28.423064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define transforms for data augmentation and normalization\ntrain_transforms = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(10),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\ntest_transforms = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])","metadata":{"execution":{"iopub.status.busy":"2023-03-22T18:31:31.531036Z","iopub.execute_input":"2023-03-22T18:31:31.531436Z","iopub.status.idle":"2023-03-22T18:31:31.53827Z","shell.execute_reply.started":"2023-03-22T18:31:31.531399Z","shell.execute_reply":"2023-03-22T18:31:31.536965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load the data\ntrain_data = datasets.ImageFolder('/kaggle/input/state-farm-distracted-driver-detection/imgs/train', transform=train_transforms)\ntest_data = datasets.ImageFolder('/kaggle/input/state-farm-distracted-driver-detection/imgs/train', transform=test_transforms)","metadata":{"execution":{"iopub.status.busy":"2023-03-22T18:38:28.747098Z","iopub.execute_input":"2023-03-22T18:38:28.747481Z","iopub.status.idle":"2023-03-22T18:38:31.788449Z","shell.execute_reply.started":"2023-03-22T18:38:28.747447Z","shell.execute_reply":"2023-03-22T18:38:31.787411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the data loaders\ntrain_loader = torch.utils.data.DataLoader(train_data, batch_size=32, shuffle=True)\ntest_loader = torch.utils.data.DataLoader(test_data, batch_size=32, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-03-22T18:38:36.464684Z","iopub.execute_input":"2023-03-22T18:38:36.465306Z","iopub.status.idle":"2023-03-22T18:38:36.470773Z","shell.execute_reply.started":"2023-03-22T18:38:36.465246Z","shell.execute_reply":"2023-03-22T18:38:36.469692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load the pre-trained ResNet model\nmodel = models.resnet18(pretrained=True)\n\n# Freeze the pre-trained layers\nfor param in model.parameters():\n    param.requires_grad = False","metadata":{"execution":{"iopub.status.busy":"2023-03-22T18:38:39.622561Z","iopub.execute_input":"2023-03-22T18:38:39.623155Z","iopub.status.idle":"2023-03-22T18:38:40.320018Z","shell.execute_reply.started":"2023-03-22T18:38:39.623105Z","shell.execute_reply":"2023-03-22T18:38:40.318919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Replace the last fully connected layer with a new one for 10 classes\nnum_ftrs = model.fc.in_features\nmodel.fc = nn.Linear(num_ftrs, 10)","metadata":{"execution":{"iopub.status.busy":"2023-03-22T18:39:03.973936Z","iopub.execute_input":"2023-03-22T18:39:03.974346Z","iopub.status.idle":"2023-03-22T18:39:03.980555Z","shell.execute_reply.started":"2023-03-22T18:39:03.974305Z","shell.execute_reply":"2023-03-22T18:39:03.97916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the loss function and optimizer\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.SGD(model.fc.parameters(), lr=0.001, momentum=0.9)\n","metadata":{"execution":{"iopub.status.busy":"2023-03-22T18:39:08.583745Z","iopub.execute_input":"2023-03-22T18:39:08.584107Z","iopub.status.idle":"2023-03-22T18:39:08.589451Z","shell.execute_reply.started":"2023-03-22T18:39:08.584075Z","shell.execute_reply":"2023-03-22T18:39:08.588297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train the model and save the best one\nbest_accuracy = 0.0\nfor epoch in range(10):\n    running_loss = 0.0\n    for inputs, labels in train_loader:\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item() * inputs.size(0)\n    epoch_loss = running_loss / len(train_data)\n    print('Epoch: {} Training Loss: {:.4f}'.format(epoch+1, epoch_loss))","metadata":{"execution":{"iopub.status.busy":"2023-03-22T18:39:11.687923Z","iopub.execute_input":"2023-03-22T18:39:11.688936Z","iopub.status.idle":"2023-03-22T18:55:06.495687Z","shell.execute_reply.started":"2023-03-22T18:39:11.688887Z","shell.execute_reply":"2023-03-22T18:55:06.493947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"  # Evaluate the model on the validation set\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for inputs, labels in test_loader:\n            outputs = model(inputs)\n            _, predicted = torch.max(outputs.data, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n    val_accuracy = 100 * correct / total\n    print('Validation Accuracy: {:.2f}%'.format(val_accuracy))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save the model if it has the best validation accuracy so far\n    if val_accuracy > best_accuracy:\n        best_accuracy = val_accuracy\n        torch.save(model.state_dict(), 'best_model.pt')\n\nprint('Best Validation Accuracy: {:.2f}%'.format(best_accuracy))","metadata":{},"execution_count":null,"outputs":[]}]}