{"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":"markdown","source":"## Random projection vs average or max pooling\n### which one is the winner?\n\nThis is a practical experiment with the ideas introduced in [this great notebook](https://www.kaggle.com/code/vldknd/linear-layer-for-dimensionality-reduction) by @vldknd\n\nI just saw that notebook and I wanted to experiment with the ideas introduced there. So, I took a simple dataset such as CIFAR10, a simple model and then implemented a simple KNN function in pytorch.\n\nLet's get started.","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-25T05:13:26.880171Z","iopub.execute_input":"2022-07-25T05:13:26.881072Z","iopub.status.idle":"2022-07-25T05:13:40.759176Z","shell.execute_reply.started":"2022-07-25T05:13:26.880968Z","shell.execute_reply":"2022-07-25T05:13:40.757709Z"}}},{"cell_type":"code","source":"!pip install timm -q","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nimport torchvision\nfrom torchvision import transforms as T\n\nfrom tqdm.autonotebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:13:53.822169Z","iopub.execute_input":"2022-07-25T05:13:53.822580Z","iopub.status.idle":"2022-07-25T05:13:56.737893Z","shell.execute_reply.started":"2022-07-25T05:13:53.822544Z","shell.execute_reply":"2022-07-25T05:13:56.736710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms = T.Compose(\n    [\n        T.Resize(224),\n        T.ToTensor(),\n        T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n\n    ]\n) ","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:13:56.739922Z","iopub.execute_input":"2022-07-25T05:13:56.740851Z","iopub.status.idle":"2022-07-25T05:13:56.750784Z","shell.execute_reply.started":"2022-07-25T05:13:56.740812Z","shell.execute_reply":"2022-07-25T05:13:56.747359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = torchvision.datasets.CIFAR10(root='.', train=False, download=True, transform=transforms)\nloader = torch.utils.data.DataLoader(dataset, batch_size=32, shuffle=True)\nmodel = timm.create_model('tf_efficientnet_b0_ns', pretrained=True, num_classes=0)\nmodel.eval()\nmodel.to('cuda');","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:13:58.497046Z","iopub.execute_input":"2022-07-25T05:13:58.497481Z","iopub.status.idle":"2022-07-25T05:14:28.004643Z","shell.execute_reply.started":"2022-07-25T05:13:58.497443Z","shell.execute_reply":"2022-07-25T05:14:28.003552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"embeds = []\nlabels = []\n\nwith torch.no_grad():\n    for images, labels_ in tqdm(loader):\n        out = model(images.to('cuda'))\n        embeds.append(out)\n        labels.append(labels_)\n\nembeds = torch.cat(embeds)\nlabels = torch.cat(labels)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:14:28.006798Z","iopub.execute_input":"2022-07-25T05:14:28.007185Z","iopub.status.idle":"2022-07-25T05:14:53.513457Z","shell.execute_reply.started":"2022-07-25T05:14:28.007144Z","shell.execute_reply":"2022-07-25T05:14:53.512554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(embeds, 'embeds.pt')\ntorch.save(labels, 'labels.pt')","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:14:56.019243Z","iopub.execute_input":"2022-07-25T05:14:56.019585Z","iopub.status.idle":"2022-07-25T05:14:56.165330Z","shell.execute_reply.started":"2022-07-25T05:14:56.019555Z","shell.execute_reply":"2022-07-25T05:14:56.164291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def normalize(a, eps=1e-8):\n    a_n = a.norm(dim=1)[:, None]\n    a_norm = a / torch.max(a_n, eps * torch.ones_like(a_n))\n    return a_norm","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:14:56.832033Z","iopub.execute_input":"2022-07-25T05:14:56.832374Z","iopub.status.idle":"2022-07-25T05:14:56.838284Z","shell.execute_reply.started":"2022-07-25T05:14:56.832346Z","shell.execute_reply":"2022-07-25T05:14:56.837168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The following function, first computes the cosine similarity of each of the embeddings to all other ones and then returns the top-k most similar embeddings for each entery. So, it does a simple KNN and returns the indices of the closest embeddings for each of the enteries in the `embeds` input tensor.","metadata":{}},{"cell_type":"code","source":"def k_nearest_neighbors(embeds, k=5):\n    normalized = normalize(embeds)\n    preds = normalized @ normalized.T\n    vals, indices = preds.sort(dim=1, descending=True)\n    \n    k += 1\n    return indices[:, 1:k].long()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:37:09.127757Z","iopub.execute_input":"2022-07-25T05:37:09.128448Z","iopub.status.idle":"2022-07-25T05:37:09.134369Z","shell.execute_reply.started":"2022-07-25T05:37:09.128414Z","shell.execute_reply":"2022-07-25T05:37:09.133198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here I first calculate a per-sample accuracy for each of the enteries (among the top-k suggestions by KNN) and then average them all. This simple KNN gets an accuracy of 83.37% with the full size embeddings (shape: (N, 1280) in efficientnet_b0 case)","metadata":{}},{"cell_type":"markdown","source":"## Full Embedding Size","metadata":{}},{"cell_type":"code","source":"preds = k_nearest_neighbors(embeds, k=5)\naccs = (labels[preds] == labels.view(-1, 1)).float().mean(dim=1)\naccs.mean()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:14:58.167083Z","iopub.execute_input":"2022-07-25T05:14:58.167423Z","iopub.status.idle":"2022-07-25T05:14:59.015495Z","shell.execute_reply.started":"2022-07-25T05:14:58.167394Z","shell.execute_reply":"2022-07-25T05:14:59.014627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Average Pooling","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:40:53.942828Z","iopub.execute_input":"2022-07-25T05:40:53.943186Z","iopub.status.idle":"2022-07-25T05:40:53.948360Z","shell.execute_reply.started":"2022-07-25T05:40:53.943150Z","shell.execute_reply":"2022-07-25T05:40:53.947187Z"}}},{"cell_type":"code","source":"embeds_avg = nn.AdaptiveAvgPool1d(64)(embeds)\n\npreds = k_nearest_neighbors(embeds_avg, k=5)\naccs = (labels[preds] == labels.view(-1, 1)).float().mean(dim=1)\naccs.mean()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:15:01.902044Z","iopub.execute_input":"2022-07-25T05:15:01.902588Z","iopub.status.idle":"2022-07-25T05:15:01.959025Z","shell.execute_reply.started":"2022-07-25T05:15:01.902537Z","shell.execute_reply":"2022-07-25T05:15:01.957985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Max Pooling","metadata":{}},{"cell_type":"code","source":"embeds_max = nn.AdaptiveMaxPool1d(64)(embeds)\n\npreds = k_nearest_neighbors(embeds_max, k=5)\naccs = (labels[preds] == labels.view(-1, 1)).float().mean(dim=1)\naccs.mean()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:15:03.471261Z","iopub.execute_input":"2022-07-25T05:15:03.471602Z","iopub.status.idle":"2022-07-25T05:15:03.523530Z","shell.execute_reply.started":"2022-07-25T05:15:03.471572Z","shell.execute_reply":"2022-07-25T05:15:03.522550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As you saw, both AvgPool1D and MaxPool1D get lower scores (which is expected) but the MaxPool performs much worse","metadata":{}},{"cell_type":"markdown","source":"## Random Projection","metadata":{}},{"cell_type":"markdown","source":"From here on, I experiment with different random projection options; from torch and sklearn.","metadata":{}},{"cell_type":"markdown","source":"### nn.Linear","metadata":{}},{"cell_type":"code","source":"all_accs = []\nfor _ in tqdm(range(100)):\n    linear = nn.Linear(1280, 64, bias=False).to('cuda')\n    embeds_rand = linear(embeds)\n\n    preds = k_nearest_neighbors(embeds_rand, k=5)\n    accs = (labels[preds] == labels.view(-1, 1)).float().mean(dim=1)\n    all_accs.append(accs.mean())\n\nall_accs = torch.stack(all_accs)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:15:04.938104Z","iopub.execute_input":"2022-07-25T05:15:04.938782Z","iopub.status.idle":"2022-07-25T05:15:09.493069Z","shell.execute_reply.started":"2022-07-25T05:15:04.938745Z","shell.execute_reply":"2022-07-25T05:15:09.492108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nprint(all_accs.mean().item())\nplt.hist(all_accs.tolist());","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:15:52.707465Z","iopub.execute_input":"2022-07-25T05:15:52.708151Z","iopub.status.idle":"2022-07-25T05:15:52.896344Z","shell.execute_reply.started":"2022-07-25T05:15:52.708112Z","shell.execute_reply":"2022-07-25T05:15:52.895543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As you see, this random projection works the best among the dim reduction techniques we saw so far. Even the lowest accuracy in this 100 test cases is higher than what Avg Pool got (0.75 > 0.73).\n\nAlso, if you are lucky enough, you can end up with an accuracy of .78 which is quite good!","metadata":{}},{"cell_type":"markdown","source":"## Simple torch.randn","metadata":{}},{"cell_type":"code","source":"all_accs = []\nfor _ in tqdm(range(100)):\n    gau = torch.randn(1280, 64).to('cuda')\n    embeds_rand = embeds @ gau\n\n    preds = k_nearest_neighbors(embeds_rand, k=5)\n    accs = (labels[preds] == labels.view(-1, 1)).float().mean(dim=1)\n    all_accs.append(accs.mean())\n\nall_accs = torch.stack(all_accs)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:16:40.153351Z","iopub.execute_input":"2022-07-25T05:16:40.153727Z","iopub.status.idle":"2022-07-25T05:16:44.637641Z","shell.execute_reply.started":"2022-07-25T05:16:40.153697Z","shell.execute_reply":"2022-07-25T05:16:44.636672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nprint(all_accs.mean().item())\nplt.hist(all_accs.tolist());","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:16:45.646073Z","iopub.execute_input":"2022-07-25T05:16:45.646738Z","iopub.status.idle":"2022-07-25T05:16:45.825425Z","shell.execute_reply.started":"2022-07-25T05:16:45.646702Z","shell.execute_reply":"2022-07-25T05:16:45.824536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## With normalizing the randn","metadata":{}},{"cell_type":"code","source":"all_accs = []\nfor _ in tqdm(range(100)):\n    gau = normalize(torch.randn(1280, 64).to('cuda'))\n    embeds_rand = embeds @ gau\n\n    preds = k_nearest_neighbors(embeds_rand, k=5)\n    accs = (labels[preds] == labels.view(-1, 1)).float().mean(dim=1)\n    all_accs.append(accs.mean())\n\nall_accs = torch.stack(all_accs)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:17:04.810052Z","iopub.execute_input":"2022-07-25T05:17:04.811113Z","iopub.status.idle":"2022-07-25T05:17:09.287390Z","shell.execute_reply.started":"2022-07-25T05:17:04.811072Z","shell.execute_reply":"2022-07-25T05:17:09.286335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nprint(all_accs.mean().item())\nplt.hist(all_accs.tolist());","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:17:09.289485Z","iopub.execute_input":"2022-07-25T05:17:09.290190Z","iopub.status.idle":"2022-07-25T05:17:09.479869Z","shell.execute_reply.started":"2022-07-25T05:17:09.290148Z","shell.execute_reply":"2022-07-25T05:17:09.478796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Using sklearn classes","metadata":{}},{"cell_type":"code","source":"from sklearn import random_projection","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:18:23.272723Z","iopub.execute_input":"2022-07-25T05:18:23.275273Z","iopub.status.idle":"2022-07-25T05:18:23.291491Z","shell.execute_reply.started":"2022-07-25T05:18:23.275234Z","shell.execute_reply":"2022-07-25T05:18:23.290500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_accs = []\nfor _ in tqdm(range(100)):\n    gau_proj = random_projection.GaussianRandomProjection(n_components=64)\n    embeds_gau = gau_proj.fit_transform(embeds.cpu().numpy())\n    embeds_gau = torch.tensor(embeds_gau).to('cuda')\n\n    preds = k_nearest_neighbors(embeds_gau, k=5)\n    accs = (labels[preds] == labels.view(-1, 1)).float().mean(dim=1)\n    all_accs.append(accs.mean())\n\nall_accs = torch.stack(all_accs)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:24:41.812237Z","iopub.execute_input":"2022-07-25T05:24:41.812588Z","iopub.status.idle":"2022-07-25T05:25:08.074260Z","shell.execute_reply.started":"2022-07-25T05:24:41.812558Z","shell.execute_reply":"2022-07-25T05:25:08.073245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nprint(all_accs.mean().item())\nplt.hist(all_accs.tolist());","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:25:10.914002Z","iopub.execute_input":"2022-07-25T05:25:10.914572Z","iopub.status.idle":"2022-07-25T05:25:11.121625Z","shell.execute_reply.started":"2022-07-25T05:25:10.914533Z","shell.execute_reply":"2022-07-25T05:25:11.120747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_accs = []\nfor _ in tqdm(range(100)):\n    spa_proj = random_projection.SparseRandomProjection(n_components=64)\n    embeds_spa = spa_proj.fit_transform(embeds.cpu().numpy())\n    embeds_spa = torch.tensor(embeds_spa).to('cuda')\n\n    preds = k_nearest_neighbors(embeds_spa, k=5)\n    accs = (labels[preds] == labels.view(-1, 1)).float().mean(dim=1)\n    all_accs.append(accs.mean())\n\nall_accs = torch.stack(all_accs)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:40:19.487915Z","iopub.execute_input":"2022-07-25T05:40:19.488292Z","iopub.status.idle":"2022-07-25T05:40:53.940927Z","shell.execute_reply.started":"2022-07-25T05:40:19.488261Z","shell.execute_reply":"2022-07-25T05:40:53.939776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nprint(all_accs.mean().item())\nplt.hist(all_accs.tolist());","metadata":{"execution":{"iopub.status.busy":"2022-07-25T05:27:19.304566Z","iopub.execute_input":"2022-07-25T05:27:19.305163Z","iopub.status.idle":"2022-07-25T05:27:19.495767Z","shell.execute_reply.started":"2022-07-25T05:27:19.305127Z","shell.execute_reply":"2022-07-25T05:27:19.494765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}