{"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":10030171,"sourceType":"datasetVersion","datasetId":6177376},{"sourceId":209940218,"sourceType":"kernelVersion"}],"dockerImageVersionId":30302,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install easyfsl","metadata":{"id":"1_tS8L5S9OTY","outputId":"5036b338-40c2-46c7-efdd-3b7d2f1118d8","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\nfrom torch import nn, optim\nfrom torch.utils.data import DataLoader\nfrom torchvision import transforms\nfrom torchvision.datasets import ImageFolder,DatasetFolder\nfrom torchvision.models import resnet18\nfrom tqdm import tqdm\n\nfrom easyfsl.samplers import TaskSampler\nfrom easyfsl.utils import plot_images, sliding_average\n\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import f1_score","metadata":{"id":"gD5jDtGZ7Krp","trusted":true,"execution":{"iopub.status.busy":"2024-11-27T14:23:48.705894Z","iopub.execute_input":"2024-11-27T14:23:48.706672Z","iopub.status.idle":"2024-11-27T14:23:49.683451Z","shell.execute_reply.started":"2024-11-27T14:23:48.706585Z","shell.execute_reply":"2024-11-27T14:23:49.682675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_size = 224\n\ntrain_set = ImageFolder(\n    root=\"/kaggle/input/driver-detection/Driver/v2_cam1_cam2_ split_by_driver/combined train/train\",\n    transform=transforms.Compose(\n        [\n            # transforms.Grayscale(num_output_channels=3),\n            transforms.RandomResizedCrop(image_size),\n            transforms.RandomHorizontalFlip(),\n            transforms.ToTensor(),\n        ]\n    ),\n)\ntest_set = ImageFolder(\n    root=\"/kaggle/input/driver-detection/Driver/v2_cam1_cam2_ split_by_driver/Camera 1/test\",\n    transform=transforms.Compose(\n        [\n            # Omniglot images have 1 channel, but our model will expect 3-channel images\n            # transforms.Grayscale(num_output_channels=3),\n            transforms.Resize([int(image_size * 1.15), int(image_size * 1.15)]),\n            transforms.CenterCrop(image_size),\n            transforms.ToTensor(),\n        ]\n    ),\n)\ntrain_set.get_labels = lambda: train_set.targets\ntest_set.idx_to_class = {idx: cl for cl, idx in test_set.class_to_idx.items()}\n","metadata":{"pycharm":{"name":"#%%\n"},"id":"OrUCQ7AslpFO","outputId":"5f78da84-6e2d-4fba-a665-de858c187a6e","_kg_hide-input":true,"_kg_hide-output":false,"trusted":true,"execution":{"iopub.status.busy":"2024-11-27T14:23:52.249308Z","iopub.execute_input":"2024-11-27T14:23:52.249837Z","iopub.status.idle":"2024-11-27T14:23:57.397651Z","shell.execute_reply.started":"2024-11-27T14:23:52.249805Z","shell.execute_reply":"2024-11-27T14:23:57.396915Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PrototypicalNetworks(nn.Module):\n    def __init__(self, backbone: nn.Module):\n        super(PrototypicalNetworks, self).__init__()\n        self.backbone = backbone\n\n    def forward(\n        self,\n        support_images: torch.Tensor,\n        support_labels: torch.Tensor,\n        query_images: torch.Tensor,\n    ) -> torch.Tensor:\n        \"\"\"\n        Predict query labels using labeled support images.\n        \"\"\"\n        # Extract the features of support and query images\n        z_support = self.backbone.forward(support_images)\n        z_query = self.backbone.forward(query_images)\n\n        # Infer the number of different classes from the labels of the support set\n        n_way = len(torch.unique(support_labels))\n        # Prototype i is the mean of all instances of features corresponding to labels == i\n        z_proto = torch.cat(\n            [\n                z_support[torch.nonzero(support_labels == label)].mean(0)\n                for label in range(n_way)\n            ]\n        )\n\n        # Compute the euclidean distance from queries to prototypes\n        dists = torch.cdist(z_query, z_proto)\n\n        # And here is the super complicated operation to transform those distances into classification scores!\n        scores = -dists\n        return scores\n\n\nconvolutional_network = resnet18(pretrained=True)\nconvolutional_network.fc = nn.Flatten()\n#print(convolutional_network)\n\nmodel = PrototypicalNetworks(convolutional_network).cuda()\n# print(model)\n","metadata":{"pycharm":{"name":"#%%\n"},"id":"iCRwLATr7Krr","outputId":"3be1f584-aaa2-4cd8-adc4-b034e97beb5b","trusted":true,"execution":{"iopub.status.busy":"2024-11-27T14:27:25.067236Z","iopub.execute_input":"2024-11-27T14:27:25.068167Z","iopub.status.idle":"2024-11-27T14:27:25.349538Z","shell.execute_reply.started":"2024-11-27T14:27:25.068131Z","shell.execute_reply":"2024-11-27T14:27:25.348430Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"N_WAY = 10  # Number of classes in a task\nN_SHOT = 30  # Number of images per class in the support set\nN_QUERY = 30  # Number of images per class in the query set\nN_EVALUATION_TASKS = 100\n\ntest_set.get_labels = lambda: test_set.targets\n# lambda: test_set.targets\n\n\ntest_sampler = TaskSampler(\n    test_set, n_way=N_WAY, n_shot=N_SHOT, n_query=N_QUERY, n_tasks=N_EVALUATION_TASKS\n)\n\ntest_loader = DataLoader(\n    test_set,\n    batch_sampler=test_sampler,\n    num_workers=4,\n    pin_memory=True,\n    collate_fn=test_sampler.episodic_collate_fn,\n)","metadata":{"pycharm":{"name":"#%%\n"},"id":"OyS0-oRV7Krt","outputId":"7b72a127-4b81-4781-8bed-07a0e54e95db","trusted":true,"execution":{"iopub.status.busy":"2024-11-27T14:24:16.616765Z","iopub.execute_input":"2024-11-27T14:24:16.617172Z","iopub.status.idle":"2024-11-27T14:24:16.624055Z","shell.execute_reply.started":"2024-11-27T14:24:16.617142Z","shell.execute_reply":"2024-11-27T14:24:16.622689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"(\n    example_support_images,\n    example_support_labels,\n    example_query_images,\n    example_query_labels,\n    example_class_ids,\n) = next(iter(test_loader))\n\nprint(example_class_ids)\nprint(test_set.idx_to_class[i] for i in list(example_class_ids))\n\nplot_images(example_support_images, \"support images\", images_per_row=N_SHOT)\nplot_images(example_query_images, \"query images\", images_per_row=N_QUERY)","metadata":{"pycharm":{"name":"#%%\n"},"id":"_FSj8NIr7Krt","outputId":"a8f63876-f1c4-4d9b-821b-ffe84baf0282","trusted":true,"execution":{"iopub.status.busy":"2024-11-27T14:24:22.105543Z","iopub.execute_input":"2024-11-27T14:24:22.106448Z","iopub.status.idle":"2024-11-27T14:25:43.726146Z","shell.execute_reply.started":"2024-11-27T14:24:22.106413Z","shell.execute_reply":"2024-11-27T14:25:43.725140Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nexample_scores = model(\n    example_support_images.cuda(),\n    example_support_labels.cuda(),\n    example_query_images.cuda(),\n).detach()\n\n_, example_predicted_labels = torch.max(example_scores.data, 1)\n\n##\n# print(example_scores.shape)\n# print(torch.softmax(example_scores, dim=1) )\n\n##\nprint(\"Some example predictions on query images:\")\nprint(\"Ground Truth - Predicted\")\nprint(\"________________________\")\nfor i in range(len(example_query_labels)):\n    print(\n        # f\"{test_set._characters[example_class_ids[example_query_labels[i]]]} / {test_set._characters[example_class_ids[example_predicted_labels[i]]]}\"\n        f\"{test_set.idx_to_class[example_class_ids[example_query_labels[i]]]} - {test_set.idx_to_class[example_class_ids[example_predicted_labels[i]]]}\"\n    )","metadata":{"pycharm":{"name":"#%%\n"},"id":"C2EhF2Fa7Kru","outputId":"6443de04-c974-4209-c634-0377aa5fcd77","trusted":true,"execution":{"iopub.status.busy":"2024-11-27T14:27:30.822372Z","iopub.execute_input":"2024-11-27T14:27:30.822751Z","iopub.status.idle":"2024-11-27T14:27:32.201644Z","shell.execute_reply.started":"2024-11-27T14:27:30.822717Z","shell.execute_reply":"2024-11-27T14:27:32.200578Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_on_one_task(\n    support_images: torch.Tensor,\n    support_labels: torch.Tensor,\n    query_images: torch.Tensor,\n) -> [int, int]:\n    \"\"\"\n    Returns the number of correct predictions of query labels, and the total number of predictions.\n    \"\"\"\n    \n    y_pred = torch.max(\n            model(support_images.cuda(), support_labels.cuda(), query_images.cuda())\n            .detach()\n            .data,\n            1,\n        )[1]\n    return y_pred\n\ndef evaluate(data_loader: DataLoader):\n    # We'll count everything and compute the ratio at the end\n    total_predictions = 0\n    correct_predictions = 0\n    f1_scores = []\n\n    # eval mode affects the behaviour of some layers (such as batch normalization or dropout)\n    # no_grad() tells torch not to keep in memory the whole computational graph (it's more lightweight this way)\n    model.eval()\n    with torch.no_grad():\n        for episode_index, (\n            support_images,\n            support_labels,\n            query_images,\n            query_labels,\n            class_ids,\n        ) in enumerate(data_loader):\n\n            y_pred = evaluate_on_one_task(\n                support_images, support_labels, query_images,\n            )\n            total_predictions += len(query_labels)\n            correct_predictions += (y_pred == query_labels.cuda()).sum().item()\n            f1_scores.append(f1_score(query_labels.cpu(), y_pred.cpu(), average='macro')) \n\n    accuracy = (correct_predictions/total_predictions)\n    avg_f1 = sum(f1_scores)/len(f1_scores)\n    print(\n        f\"Model tested on {len(data_loader)} tasks. Accuracy: {accuracy:.2%}, F1: {avg_f1}\"\n    )\n    \n    return accuracy\n\n\nevaluate(test_loader)","metadata":{"pycharm":{"name":"#%%\n"},"id":"UW5Rxifk7Kru","trusted":true,"execution":{"iopub.status.busy":"2024-11-27T14:28:00.461375Z","iopub.execute_input":"2024-11-27T14:28:00.461751Z","iopub.status.idle":"2024-11-27T14:43:14.216076Z","shell.execute_reply.started":"2024-11-27T14:28:00.461723Z","shell.execute_reply":"2024-11-27T14:43:14.214958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"N_TRAINING_EPISODES = 5000\nN_VALIDATION_TASKS = 100\n\ntrain_sampler = TaskSampler(\n    train_set, n_way=N_WAY, n_shot=N_SHOT, n_query=N_QUERY, n_tasks=N_TRAINING_EPISODES\n)\ntrain_loader = DataLoader(\n    train_set,\n    batch_sampler=train_sampler,\n    num_workers=4,\n    pin_memory=True,\n    collate_fn=train_sampler.episodic_collate_fn,\n)","metadata":{"pycharm":{"name":"#%%\n"},"id":"YW9DDxbl7Krv","outputId":"e237a732-fe67-4294-8c95-67062b0466af","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.001)\n\n\ndef fit(\n    support_images: torch.Tensor,\n    support_labels: torch.Tensor,\n    query_images: torch.Tensor,\n    query_labels: torch.Tensor,\n) -> float:\n    optimizer.zero_grad()\n    classification_scores = model(\n        support_images.cuda(), support_labels.cuda(), query_images.cuda()\n    )\n\n    loss = criterion(classification_scores, query_labels.cuda())\n    loss.backward()\n    optimizer.step()\n\n    return loss.item()","metadata":{"pycharm":{"name":"#%%\n"},"id":"0B1xX1Cb7Krv","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"log_update_frequency = 100\nacc_update_frequency = 500\n\n\nall_loss = []\nall_acc = []\n\nmodel.train()\nwith tqdm(enumerate(train_loader), total=len(train_loader)) as tqdm_train:\n    \n    for episode_index, (\n        support_images,\n        support_labels,\n        query_images,\n        query_labels,\n        _,\n    ) in tqdm_train:\n        loss_value = fit(support_images, support_labels, query_images, query_labels)\n        all_loss.append(loss_value)\n\n        if episode_index % log_update_frequency == 0:\n            tqdm_train.set_postfix(loss=sliding_average(all_loss, log_update_frequency))\n        if episode_index % acc_update_frequency == 0:\n            all_acc.append(evaluate(test_loader))\n            model.train()","metadata":{"pycharm":{"name":"#%%\n"},"id":"xQyS6uck7Krv","outputId":"ab3fe0db-aee9-42bd-e428-99ad1e7c354e","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"window_width = 50\ncumsum_vec = np.cumsum(np.insert(all_loss, 0, 0)) \nma_vec = (cumsum_vec[window_width:] - cumsum_vec[:-window_width]) / window_width\n\nplt.title(\"Loss vs # Training Episodes\")\nplt.xlabel(\"# Training Episodes\")\nplt.ylabel(\"Loss\")\nplt.plot(ma_vec, color = \"green\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"evaluate(test_loader)","metadata":{"pycharm":{"name":"#%%\n"},"id":"9bmPWd-8lpFW","outputId":"64884fde-94db-46d5-ace0-ff4eb1f2bb46","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}