{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.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":5048,"databundleVersionId":868335,"sourceType":"competition"}],"dockerImageVersionId":30262,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Understanding Data","metadata":{"id":"Y2pEdwLhD0ud"}},{"cell_type":"markdown","source":"The 10 classes to predict are:\n\n    c0: safe driving\n    c1: texting - right\n    c2: talking on the phone - right\n    c3: texting - left\n    c4: talking on the phone - left\n    c5: operating the radio\n    c6: drinking\n    c7: reaching behind\n    c8: hair and makeup\n    c9: talking to passenger\n","metadata":{"id":"NhxSi-9yYhXk"}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\npd.set_option('display.max_columns', None)\n\ndata_csv = pd.read_csv('../input/state-farm-distracted-driver-detection/driver_imgs_list.csv')\ndata_csv.columns = [x.strip() for x in data_csv.columns]\n\ndata_csv","metadata":{"id":"d-268FitPfun","outputId":"ec8d9f4b-a7dd-40ca-a84f-92aaac78d404","execution":{"iopub.status.busy":"2023-12-21T03:36:18.247332Z","iopub.execute_input":"2023-12-21T03:36:18.247898Z","iopub.status.idle":"2023-12-21T03:36:18.280738Z","shell.execute_reply.started":"2023-12-21T03:36:18.247859Z","shell.execute_reply":"2023-12-21T03:36:18.279775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import plotly.express as px\nimport plotly.graph_objects as go\nfrom skimage import io\n\ndef show_image(classname, img):\n    img = io.imread(f'../input/state-farm-distracted-driver-detection/imgs/train/{classname}/{img}')\n    fig = px.imshow(img) \n    fig.show()\n\nshow_image(data_csv.iloc[51].classname, data_csv.iloc[51].img)","metadata":{"id":"o_9enRJOP_Xb","outputId":"05752c5a-818a-48c1-e997-67a7495bf0f3","execution":{"iopub.status.busy":"2023-12-21T03:36:18.282212Z","iopub.execute_input":"2023-12-21T03:36:18.282486Z","iopub.status.idle":"2023-12-21T03:36:18.400316Z","shell.execute_reply.started":"2023-12-21T03:36:18.282460Z","shell.execute_reply":"2023-12-21T03:36:18.399456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib import pyplot as plt\nfrom matplotlib import animation\nfrom PIL import Image\nimport matplotlib.image as mpimg\nimport matplotlib.animation as animation\n\ndef animate_subject(snapshots, fps=30, nSeconds=15, anim_title=\"new_animation\"):\n\n    # First set up the figure, the axis, and the plot element we want to animate\n    fig = plt.figure( figsize=(8,8) )\n\n    a = snapshots[0]\n    im = plt.imshow(a, interpolation='none', aspect='auto', vmin=0, vmax=1)\n\n    def animate_func(i):\n        if i % fps == 0:\n            print( '.', end ='' )\n\n        im.set_array(snapshots[i])\n        return [im]\n\n    anim = animation.FuncAnimation(\n                                fig, \n                                animate_func, \n                                frames = nSeconds * fps,\n                                interval = 1000/fps, # in ms\n                                )\n\n    anim.save(f'{anim_title}.mp4', fps=fps, extra_args=['-vcodec', 'libx264'])\n    plt.clf()\n\n    from IPython.display import Video\n\n    return Video(f\"{anim_title}.mp4\", embed=True)","metadata":{"id":"6Zlchl_IBHUv","execution":{"iopub.status.busy":"2023-12-21T03:36:18.401649Z","iopub.execute_input":"2023-12-21T03:36:18.401996Z","iopub.status.idle":"2023-12-21T03:36:18.414429Z","shell.execute_reply.started":"2023-12-21T03:36:18.401963Z","shell.execute_reply":"2023-12-21T03:36:18.413361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = data_csv[data_csv.subject == 'p081']\nsnapshots = [ mpimg.imread(f'../input/state-farm-distracted-driver-detection/imgs/train/{df.iloc[x].classname}/{df.iloc[x].img}') for x in range(len(df))]\nanimate_subject(snapshots)","metadata":{"id":"U5ofHcI-RRl5","outputId":"f30204f7-4fd9-4204-c5be-3743e63d2467","execution":{"iopub.status.busy":"2023-12-21T03:36:18.417140Z","iopub.execute_input":"2023-12-21T03:36:18.417657Z","iopub.status.idle":"2023-12-21T03:36:52.822257Z","shell.execute_reply.started":"2023-12-21T03:36:18.417619Z","shell.execute_reply":"2023-12-21T03:36:52.821184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from plotly.subplots import make_subplots\n\nsubs = make_subplots(rows=1, cols=2, subplot_titles=[\"Subject\", \"Classname\"])\n\nsubs.add_trace(go.Histogram(x=data_csv['subject']), row=1, col=1)\nsubs.add_trace(go.Histogram(x=data_csv['classname']), row=1, col=2)\n\nsubs.show()","metadata":{"id":"cAg-VWppWvfw","outputId":"354a16fa-61c9-40d2-c0bb-b2aeae421c96","execution":{"iopub.status.busy":"2023-12-21T03:36:52.823659Z","iopub.execute_input":"2023-12-21T03:36:52.824000Z","iopub.status.idle":"2023-12-21T03:36:53.077823Z","shell.execute_reply.started":"2023-12-21T03:36:52.823965Z","shell.execute_reply":"2023-12-21T03:36:53.076837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preparing Data for Training/Validation","metadata":{"id":"aeiDzvWtD62W"}},{"cell_type":"code","source":"import torch\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"Using {device} device\")","metadata":{"id":"BNdxtvcL-aAN","outputId":"7afc8488-c9e0-43b3-8193-7d7366487f77","execution":{"iopub.status.busy":"2023-12-21T03:36:53.079142Z","iopub.execute_input":"2023-12-21T03:36:53.079461Z","iopub.status.idle":"2023-12-21T03:36:53.085409Z","shell.execute_reply.started":"2023-12-21T03:36:53.079431Z","shell.execute_reply":"2023-12-21T03:36:53.084225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import nn\nfrom torch.utils.data import Dataset, DataLoader, random_split \nimport torchvision.transforms as T\n\nclass state_farm_dataset(Dataset):\n    def __init__(self, classes , images, transform=None, target_transform=None):\n        self.labels = classes\n        self.images = images \n        self.transform = transform\n        self.target_transform = target_transform\n\n    def __len__(self):\n        return len(self.labels)\n    \n    def __getitem__(self, idx):\n        img = Image.open(f'../input/state-farm-distracted-driver-detection/imgs/train/{self.labels[idx]}/{self.images[idx]}')\n        label = self.labels[idx]\n\n        if self.transform: img = self.transform(img)\n        if self.target_transform: label = self.target_transform(label)\n\n        return img, label\n","metadata":{"id":"dTrePGLFbK4r","execution":{"iopub.status.busy":"2023-12-21T03:36:53.086771Z","iopub.execute_input":"2023-12-21T03:36:53.087094Z","iopub.status.idle":"2023-12-21T03:36:53.097630Z","shell.execute_reply.started":"2023-12-21T03:36:53.087065Z","shell.execute_reply":"2023-12-21T03:36:53.096604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def encode_labels(y):\n    idx = int(y.split('c')[1])\n    return idx","metadata":{"id":"u0QD5uQkvqm2","execution":{"iopub.status.busy":"2023-12-21T03:36:53.098826Z","iopub.execute_input":"2023-12-21T03:36:53.099110Z","iopub.status.idle":"2023-12-21T03:36:53.109491Z","shell.execute_reply.started":"2023-12-21T03:36:53.099084Z","shell.execute_reply":"2023-12-21T03:36:53.108524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def image_transform(x):\n    x= T.Resize(200)(x)\n    x =T.Grayscale()(x)\n    return T.ToTensor()(x)\n","metadata":{"id":"hg-pcw5jGRvf","execution":{"iopub.status.busy":"2023-12-21T03:36:53.110622Z","iopub.execute_input":"2023-12-21T03:36:53.110938Z","iopub.status.idle":"2023-12-21T03:36:53.118251Z","shell.execute_reply.started":"2023-12-21T03:36:53.110911Z","shell.execute_reply":"2023-12-21T03:36:53.117102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size=64\ndataset = state_farm_dataset(data_csv.classname, data_csv.img,target_transform=encode_labels, transform=image_transform)\nsemi_train_data, train_data, test_data = random_split(dataset, [1000, 15000, 6424])\nsemi_train_dataloader = DataLoader(semi_train_data, batch_size=batch_size, shuffle=True)\ntrain_dataloader = DataLoader(train_data, batch_size=batch_size, shuffle=True)\ntest_dataloader = DataLoader(test_data, batch_size=batch_size, shuffle=True)","metadata":{"id":"EaLD2RAHhGCq","execution":{"iopub.status.busy":"2023-12-21T03:36:53.122841Z","iopub.execute_input":"2023-12-21T03:36:53.123143Z","iopub.status.idle":"2023-12-21T03:36:53.133938Z","shell.execute_reply.started":"2023-12-21T03:36:53.123117Z","shell.execute_reply":"2023-12-21T03:36:53.132866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_img = Image.open('../input/state-farm-distracted-driver-detection/imgs/test/img_1.jpg')\nimage_transform(test_img).shape","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:36:53.135317Z","iopub.execute_input":"2023-12-21T03:36:53.135668Z","iopub.status.idle":"2023-12-21T03:36:53.156525Z","shell.execute_reply.started":"2023-12-21T03:36:53.135633Z","shell.execute_reply":"2023-12-21T03:36:53.155507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Image Sample of individual classes to predict","metadata":{"id":"_Bcp43x-EGhb"}},{"cell_type":"code","source":"import torchvision.utils as U\nsample_imgs, sample_labels = next(iter(semi_train_dataloader))\n\nclasses = torch.Tensor([]) \n\nfor x in range(10):\n    idx = (sample_labels == x).nonzero(as_tuple=False)[0][0]\n    classes = torch.cat((classes, sample_imgs[idx]), 0)\n\ngrid = U.make_grid(classes.unsqueeze(dim=1), nrows=8)\nplt.figure(figsize=(25, 25))\nplt.imshow(grid.permute(1, 2, 0))","metadata":{"id":"pRMkAEEPxFJC","outputId":"b4369368-6b5c-48a7-a92d-341ff9ff9b04","execution":{"iopub.status.busy":"2023-12-21T03:36:53.157703Z","iopub.execute_input":"2023-12-21T03:36:53.158011Z","iopub.status.idle":"2023-12-21T03:36:54.828187Z","shell.execute_reply.started":"2023-12-21T03:36:53.157978Z","shell.execute_reply":"2023-12-21T03:36:54.827194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(semi_train_dataloader.dataset))\nsample_imgs.shape","metadata":{"id":"46lyhYUFSdEP","outputId":"fcaa0949-cdff-4241-e43f-5b0e59d01f9d","execution":{"iopub.status.busy":"2023-12-21T03:36:54.829574Z","iopub.execute_input":"2023-12-21T03:36:54.829900Z","iopub.status.idle":"2023-12-21T03:36:54.837739Z","shell.execute_reply.started":"2023-12-21T03:36:54.829869Z","shell.execute_reply":"2023-12-21T03:36:54.836467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Defining Network and functions","metadata":{"id":"TyDlPDe0ERDC"}},{"cell_type":"code","source":"class BasicNetwork(nn.Module):\n    def __init__(self):\n        super(BasicNetwork, self).__init__()\n        self.non_linearity = nn.Tanh\n        self.max_pooling_size = 2\n\n        self.conv_stack = nn.Sequential(\n            nn.Conv2d(1, 3,(7,7), stride=3 ),\n            self.non_linearity(),\n\n            nn.Conv2d(3, 6, (3, 3), stride=2),\n            self.non_linearity(),\n\n            nn.Conv2d(6, 12, (2, 2)),\n            self.non_linearity(),\n\n            nn.MaxPool2d(self.max_pooling_size),\n        )\n\n        self.linear_stack = nn.Sequential(\n            nn.Linear(3780, 10),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        conv_output = self.conv_stack(x)\n        # print(\"Convolution output: \", conv_output.shape, conv_output.flatten(start_dim=1).shape)\n        linear_output = self.linear_stack(conv_output.flatten(start_dim=1))\n        return linear_output\n        ","metadata":{"id":"isbFzxf2SZM8","execution":{"iopub.status.busy":"2023-12-21T03:36:54.839363Z","iopub.execute_input":"2023-12-21T03:36:54.839705Z","iopub.status.idle":"2023-12-21T03:36:54.849703Z","shell.execute_reply.started":"2023-12-21T03:36:54.839662Z","shell.execute_reply":"2023-12-21T03:36:54.848702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_net = BasicNetwork()\ntest_net(sample_imgs).argmax(dim=1)","metadata":{"id":"R5pdfSJqTcek","outputId":"14e39575-1974-4302-ce7d-ea68e17dc430","execution":{"iopub.status.busy":"2023-12-21T03:36:54.851069Z","iopub.execute_input":"2023-12-21T03:36:54.851419Z","iopub.status.idle":"2023-12-21T03:36:54.878530Z","shell.execute_reply.started":"2023-12-21T03:36:54.851384Z","shell.execute_reply":"2023-12-21T03:36:54.877575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Function Definations","metadata":{"id":"wJJmryDhvC1D"}},{"cell_type":"code","source":"def train(model, dataloader, loss_func, optimizer_func, verbose=True):\n    model.train()\n    loss_arr = []\n\n    for batch, (X, y) in enumerate(dataloader):\n        X,y = X.to(device), y.to(device)\n        pred = model(X)\n        loss = loss_func(pred, y)\n        optimizer_func.zero_grad()\n        loss.backward()\n        optimizer_func.step()\n\n        loss_arr.append(loss.item())\n        \n        if(verbose):\n            if(batch%100 == 0):\n                print(f'batch: {batch} loss: {loss.item()} [{batch*len(X)} / {len(dataloader.dataset)}]')\n    return loss_arr","metadata":{"id":"ZSXR060lbyzg","execution":{"iopub.status.busy":"2023-12-21T03:36:54.879690Z","iopub.execute_input":"2023-12-21T03:36:54.880173Z","iopub.status.idle":"2023-12-21T03:36:54.888710Z","shell.execute_reply.started":"2023-12-21T03:36:54.880142Z","shell.execute_reply":"2023-12-21T03:36:54.887704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef test(model, dataloader):\n    correct = 0\n    total = len(dataloader.dataset)\n\n    for batch, (X, y) in enumerate(dataloader):\n        X, y = X.to(device), y.to(device)\n        pred = model(X)\n        correct += sum(pred.argmax(dim=1) == y)\n    return [correct, total]","metadata":{"id":"D8yCyDaLcwto","execution":{"iopub.status.busy":"2023-12-21T03:36:54.890003Z","iopub.execute_input":"2023-12-21T03:36:54.890372Z","iopub.status.idle":"2023-12-21T03:36:54.897846Z","shell.execute_reply.started":"2023-12-21T03:36:54.890335Z","shell.execute_reply":"2023-12-21T03:36:54.897003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision.utils as U\n\ndef print_filters(image, layer, nrows=8, show_input_image=True):\n    if show_input_image: plt.imshow(image.squeeze(dim=0), cmap=\"gray\")\n    image = image.to(device)\n    output = layer(image)\n    grid = U.make_grid(output.unsqueeze(dim=1), nrow=output.shape[0]//nrows+1)\n    return grid\n","metadata":{"id":"_fOZSvKlvSaO","execution":{"iopub.status.busy":"2023-12-21T03:36:54.899153Z","iopub.execute_input":"2023-12-21T03:36:54.899415Z","iopub.status.idle":"2023-12-21T03:36:54.907775Z","shell.execute_reply.started":"2023-12-21T03:36:54.899391Z","shell.execute_reply":"2023-12-21T03:36:54.906918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# More tests on semi_train_dataloader ","metadata":{"id":"JnW2QToPb1Gl"}},{"cell_type":"markdown","source":"### Testing Multiple Learning Rates on the model","metadata":{"id":"RuUIhd5DEwWO"}},{"cell_type":"code","source":"def test_learning_rates(Network, dataloader, learning_rates):\n    losses = []\n    for lr in learning_rates:\n        loss_func = nn.CrossEntropyLoss()\n        main_net = Network().to(device)\n        optimizer_func = torch.optim.SGD(main_net.parameters(), lr=lr)\n        \n        loss_arr = train(main_net, dataloader, loss_func, optimizer_func, verbose=False)\n\n        losses.append(loss_arr)\n    return losses","metadata":{"id":"EEOF7SxbdUAN","execution":{"iopub.status.busy":"2023-12-21T03:36:54.908885Z","iopub.execute_input":"2023-12-21T03:36:54.909151Z","iopub.status.idle":"2023-12-21T03:36:54.920697Z","shell.execute_reply.started":"2023-12-21T03:36:54.909125Z","shell.execute_reply":"2023-12-21T03:36:54.919786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import plotly.express as px\nfrom plotly.subplots import make_subplots\nimport plotly.graph_objects as go\n\nlearning_rates = [0.0001, 0.001, 0.01, 0.1, 0.9]\ntotal_losses = test_learning_rates(BasicNetwork, semi_train_dataloader, learning_rates)\n\nrows,cols, idx = 2, 3, 0\nsubs = make_subplots(rows=rows, cols=cols, subplot_titles=learning_rates)\n\nfor row in range(rows):\n    for col in range(cols):\n        if(row == 1 and col == 2): break;\n        subs.add_trace(go.Scatter(y=total_losses[idx]), row=row+1, col=col+1)\n        idx+=1\n\nsubs.show()","metadata":{"id":"qRMkAGbmdlCO","execution":{"iopub.status.busy":"2023-12-21T03:36:54.921838Z","iopub.execute_input":"2023-12-21T03:36:54.922195Z","iopub.status.idle":"2023-12-21T03:37:43.217101Z","shell.execute_reply.started":"2023-12-21T03:36:54.922158Z","shell.execute_reply":"2023-12-21T03:37:43.216096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Testing Model performance on epochs. To Figure out how long should the model be trained","metadata":{"id":"moG2M1NNE0VC"}},{"cell_type":"code","source":"def test_epochs(Network, dataloader, number_of_epochs):\n    epoch_list = []\n    accuracy_list = []\n    main_net = Network().to(device)\n    loss_func = nn.CrossEntropyLoss()\n    optimizer_func = torch.optim.SGD(main_net.parameters(), lr=0.9)\n\n    for i in range(number_of_epochs):\n        epoch_list.append(i)\n        train(main_net, dataloader, loss_func, optimizer_func, verbose=False)\n        correct, total = test(main_net, dataloader)\n        accuracy_list.append(correct/total)\n    return accuracy_list\n","metadata":{"id":"ARymBobrd7fr","execution":{"iopub.status.busy":"2023-12-21T03:37:43.218745Z","iopub.execute_input":"2023-12-21T03:37:43.219173Z","iopub.status.idle":"2023-12-21T03:37:43.226967Z","shell.execute_reply.started":"2023-12-21T03:37:43.219135Z","shell.execute_reply":"2023-12-21T03:37:43.226015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epoch_test_result = test_epochs(BasicNetwork, semi_train_dataloader, 20)","metadata":{"id":"blMfFKPWeUup","execution":{"iopub.status.busy":"2023-12-21T03:37:43.228095Z","iopub.execute_input":"2023-12-21T03:37:43.228403Z","iopub.status.idle":"2023-12-21T03:43:12.052060Z","shell.execute_reply.started":"2023-12-21T03:37:43.228356Z","shell.execute_reply":"2023-12-21T03:43:12.051041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epoch_test_result = [x.cpu() for x in epoch_test_result]\npx.line(epoch_test_result)","metadata":{"id":"TQT3Gb0JhD_A","execution":{"iopub.status.busy":"2023-12-21T03:43:12.053576Z","iopub.execute_input":"2023-12-21T03:43:12.053979Z","iopub.status.idle":"2023-12-21T03:43:12.162024Z","shell.execute_reply.started":"2023-12-21T03:43:12.053939Z","shell.execute_reply":"2023-12-21T03:43:12.161140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Actual Model Training\n\n","metadata":{"id":"vrF_9MH9eie9"}},{"cell_type":"code","source":"myNet = BasicNetwork()\nmyNet.to(device)","metadata":{"id":"glqvVmL0xGYO","outputId":"1b16c2cc-86b3-40a0-a370-d35334497006","execution":{"iopub.status.busy":"2023-12-21T03:43:12.163373Z","iopub.execute_input":"2023-12-21T03:43:12.163990Z","iopub.status.idle":"2023-12-21T03:43:12.172650Z","shell.execute_reply.started":"2023-12-21T03:43:12.163957Z","shell.execute_reply":"2023-12-21T03:43:12.171677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p losses  ","metadata":{"id":"zsDw8cTt4_Kd","outputId":"4037c141-fe98-495d-ccb2-9dd0ff1df854","execution":{"iopub.status.busy":"2023-12-21T03:43:12.174037Z","iopub.execute_input":"2023-12-21T03:43:12.174314Z","iopub.status.idle":"2023-12-21T03:43:13.209933Z","shell.execute_reply.started":"2023-12-21T03:43:12.174288Z","shell.execute_reply":"2023-12-21T03:43:13.208657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = 22\nloss_func = nn.CrossEntropyLoss()\noptimizer= torch.optim.SGD(myNet.parameters(), lr=0.9)\n\nconv1_snaps = [] #0\nconv2_snaps = [] #2\nconv3_snaps = [] #4\n\ncounter = 0\nfor i in range(epochs):\n    print(\"Epoch: \", i+1)\n    loss_arr = train(myNet, train_dataloader, loss_func, optimizer)\n    torch.save(loss_arr, f'losses/{counter}.txt') \n    counter += 1\n\n    conv1_output = myNet.conv_stack[0](sample_imgs[0].to(device))\n    conv2_output = myNet.conv_stack[2](conv1_output)\n\n    conv1_snaps.append(print_filters(sample_imgs[0], myNet.conv_stack[0], show_input_image=False))\n    conv2_snaps.append(print_filters(conv1_output, myNet.conv_stack[2], show_input_image=False))\n    conv3_snaps.append(print_filters(conv2_output, myNet.conv_stack[4], show_input_image=False))","metadata":{"id":"_oq6gH1ZBs2K","outputId":"e90b80dc-6cc5-4e80-eb7e-08c9a3ec713f","execution":{"iopub.status.busy":"2023-12-21T03:43:13.212020Z","iopub.execute_input":"2023-12-21T03:43:13.212449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"[correct, total ] = test(myNet, test_dataloader)\nprint(\"Accuracy: \", correct/total)","metadata":{"id":"vELBhTqwARPo","outputId":"3c2d4eef-8bc5-48b8-970c-ee367f8eaf9a","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Result Visualization","metadata":{"id":"Eze7zeN2E_4u"}},{"cell_type":"markdown","source":"### Loss Graph","metadata":{"id":"HqdRaXMuFB42"}},{"cell_type":"code","source":"total_loss = torch.Tensor([])\nfor i in range(10):\n    temp = torch.Tensor(torch.load(f'losses/{i}.txt'))\n    total_loss = torch.cat((total_loss, temp), 0)\n\npx.line(y=total_loss)","metadata":{"id":"_08Glx-EBkmn","outputId":"b1eb692c-c0ca-433f-dc2e-d4c666ff9fc1","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Confusion Matrix","metadata":{"id":"QG4b-pVpFGRA"}},{"cell_type":"code","source":"import plotly.express as px\n\ndef build_confusion_matrix(main_net, dataloader, classes):\n    total_preds = torch.Tensor([]).long().to(device)\n    total_labels = torch.Tensor([]).long().to(device)\n    \n    for batch,(X, y) in enumerate(dataloader):\n        X, y = X.to(device), y.long().to(device)\n        pred = main_net(X).long()\n\n        total_labels = torch.cat((total_labels, y), 0)\n        total_preds = torch.cat((total_preds, pred.argmax(dim=1)), 0)\n\n    matrix = torch.Tensor([[0, 0, 0, 0, 0, 0, 0, 0, 0, 0] for x in range(classes)]).long().to(device)\n\n    for idx, label in enumerate(total_labels):\n        matrix[label][total_preds[idx]] += 1\n\n    return matrix.cpu()\n\nconf_matrix = build_confusion_matrix(myNet, test_dataloader, classes = 10)\npx.imshow(conf_matrix, text_auto=True)","metadata":{"id":"D8CS9J8aFFDW","outputId":"574712fe-4fc3-4775-dba8-119d11caa8b7","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualizing how filters changed during training","metadata":{"id":"DUxrKlLNFPmm"}},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")\n\n#First Convolutional layer\ngrid_arr =  [U.make_grid(conv1_snaps[i].unsqueeze(dim=1), nrow=8).cpu().permute(1, 2, 0) for i in range(epochs)]\nanimate_subject(grid_arr, fps=1, nSeconds=10, anim_title=\"conv1\")","metadata":{"id":"ca_as3Ca1QHQ","outputId":"0955ee2d-9096-4f5e-92bf-9c17ab641cda","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Second Convolutional Layer\n\ngrid_arr =  [U.make_grid(conv2_snaps[i].unsqueeze(dim=1), nrow=8).cpu().permute(1, 2, 0) for i in range(epochs)]\nanimate_subject(grid_arr, fps=1, nSeconds=10, anim_title=\"conv2\")","metadata":{"id":"rVuEaA00-Gl0","outputId":"3d16d5fe-5c5f-4bab-8528-92538c2abc31","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Third Convolutional Layer\ngrid_arr =  [U.make_grid(conv3_snaps[i].unsqueeze(dim=1), nrow=8).cpu().permute(1, 2, 0) for i in range(epochs)]\nanimate_subject(grid_arr, fps=1, nSeconds=10, anim_title=\"conv3\")","metadata":{"id":"XbXqqWl3_xRV","outputId":"6cadd55b-c205-4da5-cd74-c68a06d62421","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\n\ndirectory ='../input/state-farm-distracted-driver-detection/imgs/test'\n\ncolumn_names=['img', 'c0', 'c1', 'c2', 'c3', 'c4', 'c5', 'c6', 'c7', 'c8', 'c9']\narr = np.array([['img', 'c0', 'c1', 'c2', 'c3', 'c4', 'c5', 'c6', 'c7', 'c8', 'c9']])\ntotal = len(os.listdir(directory))\n\nsig = nn.Sigmoid()\n\nfor idx, filename in enumerate(os.listdir(directory)):\n    img = image_transform(Image.open(f'{directory}/{filename}')).to(device)\n    pred = myNet(img.unsqueeze(dim=0))\n    pred = sig(pred)\n    temp = np.append(filename, pred.cpu().detach().numpy())\n    arr = np.append(arr, temp.reshape(1, -1), axis=0)\n    \n    if(idx % 1000 == 0):\n        print(f'validation progress... [{idx}/{total}] {idx/total*100}')\n\nsubmission = pd.DataFrame(arr,columns=column_names)[1:]\nsubmission = submission.astype(dtype={\n    'c0': float,\n    'c1': float,\n    'c2': float,\n    'c3': float,\n    'c4': float,\n    'c5': float,\n    'c6': float,\n    'c7': float,\n    'c8': float,\n    'c9': float,\n})\nsubmission.to_csv('submission.csv', index=False)\nsubmission","metadata":{"id":"v48-q3ReH6VR","trusted":true},"execution_count":null,"outputs":[]}]}