{"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":"# 1 - Summary\n#### Task: Multi-Class Classification\n#### N° Classes: 3\n#### Class Names & Indices: \"astro\" (0), \"cort\" (1), \"shsy5y\", (2)\n#### N° Samples: 2578\n#### Imbalanced: slightly\n#### EDA: [here](https://www.kaggle.com/code/mahmoudlimam/sartorius-eda-for-classification)\n#### Original Image Size: 520x704\n#### Resized to: 512x512\n#### Pre-Processing: Histogram Equalization (only used for the tabular task)\n## Deep Learning Approach:\n#### N° Validation Samples: 200\n#### Architecture: CNN\n#### Model Size: 547,427 parameters\n#### Transfer Learning: No\n#### Batch Size: 64\n#### N° Epochs: 10\n#### Learning Rate: 0.001\n#### Mixed Precision Training: Yes (fp16)\n#### Gradient Accumulation: No (1 step)\n#### Metrics: Accuracy & F1-Score (to account for the imbalance)\n#### Main Libraries: PyTorch, TorchVision, Accelerate\n#### GPU: P100\n#### Results:\n* **Over 95% performance on training & validation data.**\n* **Equalization didn't improve training, it just slowed it down. Thus i didn't actually use it in training.**\n* **The slight class imbalance didn't affect performance; Got practically equal values of Accuracy & F1-Score.**\n* **256x256 images resulted in awful results (around 44% on both metrics), 512x512 showed a 50% improvement.**  \n* **Note: 256x256 training/results are not shown in the notebook**\n\n## Tabular ML Approach (RF with Manually Selected Features):\n#### N° Validation Samples: 238\n#### Results:\n* **100% performance on training data**\n* **88% performance on testing data**\n* **Quite Overfit**\n* **No imbalance issues**","metadata":{}},{"cell_type":"markdown","source":"# 2 - Setup","metadata":{}},{"cell_type":"code","source":"!pip install torchsummary","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:56:55.161840Z","iopub.execute_input":"2022-12-11T10:56:55.162281Z","iopub.status.idle":"2022-12-11T10:57:04.702907Z","shell.execute_reply.started":"2022-12-11T10:56:55.162243Z","shell.execute_reply":"2022-12-11T10:57:04.701691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport plotly.express as px\nimport os\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision\nfrom torchvision import transforms as T\nfrom tqdm import tqdm\nimport tqdm.notebook as tq\nfrom PIL import Image\nimport random\nfrom torchsummary import summary\nfrom accelerate import Accelerator\nfrom torchmetrics import Accuracy, F1Score\nimport umap\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.ensemble import RandomForestClassifier\nfrom sklearn.metrics import classification_report","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-11T10:57:04.706975Z","iopub.execute_input":"2022-12-11T10:57:04.707279Z","iopub.status.idle":"2022-12-11T10:57:04.714576Z","shell.execute_reply.started":"2022-12-11T10:57:04.707248Z","shell.execute_reply":"2022-12-11T10:57:04.713595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"HEIGHT = 512\nWIDTH = 512\nLABEL_MAP = {\"astro\":0, \"cort\":1, \"shsy5y\":2}\nBATCH_SIZE = 64\nNEPOCHS = 10\nGRADIENT_ACCUMULATION_STEPS = 1\nMIXED_PRECISION = \"fp16\"\nLEARNING_RATE = 0.001\nBEST_MODEL_PATH = \"model.pt\"","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:04.716088Z","iopub.execute_input":"2022-12-11T10:57:04.716417Z","iopub.status.idle":"2022-12-11T10:57:04.724806Z","shell.execute_reply.started":"2022-12-11T10:57:04.716378Z","shell.execute_reply":"2022-12-11T10:57:04.723898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.__version__","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:04.728764Z","iopub.execute_input":"2022-12-11T10:57:04.729316Z","iopub.status.idle":"2022-12-11T10:57:04.738934Z","shell.execute_reply.started":"2022-12-11T10:57:04.729290Z","shell.execute_reply":"2022-12-11T10:57:04.738021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3 - Data Wrangling & Preparation","metadata":{}},{"cell_type":"code","source":"data1 = pd.read_csv(\"../input/sartorius-cell-instance-segmentation/train.csv\")\ndata1_images = \"../input/sartorius-cell-instance-segmentation/train\"\nsemi_dir = \"../input/sartorius-cell-instance-segmentation/train_semi_supervised\"","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:04.742031Z","iopub.execute_input":"2022-12-11T10:57:04.742291Z","iopub.status.idle":"2022-12-11T10:57:05.020543Z","shell.execute_reply.started":"2022-12-11T10:57:04.742257Z","shell.execute_reply":"2022-12-11T10:57:05.019586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def transform1(d, path):\n    this = d.groupby(\"id\")[\"cell_type\"].unique()\n    arr = np.empty(shape=(len(this),2),dtype=\"object\")\n    for i in range(len(this)):\n        arr[i,1] = this[i][0]\n        arr[i,0] = path + \"/\" + this.index[i] + \".png\"\n    return pd.DataFrame(arr, columns=(\"path\",\"class\"))","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.022135Z","iopub.execute_input":"2022-12-11T10:57:05.022518Z","iopub.status.idle":"2022-12-11T10:57:05.029343Z","shell.execute_reply.started":"2022-12-11T10:57:05.022482Z","shell.execute_reply":"2022-12-11T10:57:05.028274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def transform2(path):\n    liss = os.listdir(path)\n    arr = np.empty(shape=(len(liss),2),dtype=\"object\")\n    for i in range(len(liss)):\n        arr[i,0] = os.path.join(path,liss[i])\n        v = liss[i].split(\"[\")[0]\n        if v==\"astros\":\n            v=\"astro\"\n        arr[i,1] = v\n    return pd.DataFrame(arr, columns=(\"path\",\"class\"))","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.030732Z","iopub.execute_input":"2022-12-11T10:57:05.031347Z","iopub.status.idle":"2022-12-11T10:57:05.043761Z","shell.execute_reply.started":"2022-12-11T10:57:05.031309Z","shell.execute_reply":"2022-12-11T10:57:05.042861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data1 = transform1(data1, data1_images)\ndata2 = transform2(semi_dir)\ndata = pd.concat([data1,data2]).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.044961Z","iopub.execute_input":"2022-12-11T10:57:05.045322Z","iopub.status.idle":"2022-12-11T10:57:05.109571Z","shell.execute_reply.started":"2022-12-11T10:57:05.045280Z","shell.execute_reply":"2022-12-11T10:57:05.108754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.110708Z","iopub.execute_input":"2022-12-11T10:57:05.112201Z","iopub.status.idle":"2022-12-11T10:57:05.124448Z","shell.execute_reply.started":"2022-12-11T10:57:05.112165Z","shell.execute_reply":"2022-12-11T10:57:05.123413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_image(img):\n    return T.functional.equalize(img)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.129055Z","iopub.execute_input":"2022-12-11T10:57:05.129313Z","iopub.status.idle":"2022-12-11T10:57:05.135187Z","shell.execute_reply.started":"2022-12-11T10:57:05.129289Z","shell.execute_reply":"2022-12-11T10:57:05.133828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class image_dataset(Dataset):\n    \n    def __init__(self, df, processing_function=None):\n        super().__init__()\n        self.df = df\n        self.processing_function = processing_function\n        \n        \n    def __len__(self):\n        return self.df.shape[0]\n    \n    def __getitem__(self, ind):\n        path, label = tuple(self.df.iloc[ind,:])\n        label = torch.as_tensor(LABEL_MAP[label], dtype=torch.float)\n        img = Image.open(path)\n        if self.processing_function!=None:\n            img = self.processing_function(img)\n        img = T.Resize(size=(HEIGHT,WIDTH))(img)\n        img = T.PILToTensor()(img)/255.0\n        return img, label","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.136953Z","iopub.execute_input":"2022-12-11T10:57:05.137284Z","iopub.status.idle":"2022-12-11T10:57:05.146170Z","shell.execute_reply.started":"2022-12-11T10:57:05.137251Z","shell.execute_reply":"2022-12-11T10:57:05.145283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def split(df, val_size=200):\n    val_samples = data.sample(n=val_size)\n    train_samples = df.drop(index=val_samples.index)\n    return train_samples, val_samples","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.147664Z","iopub.execute_input":"2022-12-11T10:57:05.148031Z","iopub.status.idle":"2022-12-11T10:57:05.155960Z","shell.execute_reply.started":"2022-12-11T10:57:05.147998Z","shell.execute_reply":"2022-12-11T10:57:05.155117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"traindata, valdata = split(data)\ntraindata = image_dataset(traindata, processing_function=None)\nvaldata = image_dataset(valdata, processing_function=None)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.158893Z","iopub.execute_input":"2022-12-11T10:57:05.159243Z","iopub.status.idle":"2022-12-11T10:57:05.167025Z","shell.execute_reply.started":"2022-12-11T10:57:05.159218Z","shell.execute_reply":"2022-12-11T10:57:05.166081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(traindata, batch_size=BATCH_SIZE, shuffle=True)\nval_loader = DataLoader(valdata, batch_size=BATCH_SIZE, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.168652Z","iopub.execute_input":"2022-12-11T10:57:05.169010Z","iopub.status.idle":"2022-12-11T10:57:05.175629Z","shell.execute_reply.started":"2022-12-11T10:57:05.168976Z","shell.execute_reply":"2022-12-11T10:57:05.174691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_examples(n):\n    samples = data.sample(n=n, random_state=11)\n    for i in range(n):\n        fig, axes = plt.subplots(nrows=1, ncols=2, figsize=(20,10))\n        path, label = samples.iloc[i,0], samples.iloc[i,1]\n        img = Image.open(path)\n        eq_img = process_image(img)\n        img = torchvision.transforms.Resize(size=(HEIGHT,WIDTH))(img)\n        eq_img = torchvision.transforms.Resize(size=(HEIGHT,WIDTH))(eq_img)\n        axes[0].imshow(img, cmap=\"gray\")\n        axes[0].set_title(f\"Original\", fontsize=12)\n        axes[1].imshow(eq_img, cmap=\"gray\")\n        axes[1].set_title(f\"Equalized\", fontsize=12)\n        plt.suptitle(f'\"{label}\" Sample', fontsize=20)\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.177177Z","iopub.execute_input":"2022-12-11T10:57:05.177596Z","iopub.status.idle":"2022-12-11T10:57:05.188187Z","shell.execute_reply.started":"2022-12-11T10:57:05.177562Z","shell.execute_reply":"2022-12-11T10:57:05.187285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_examples(10)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:05.191203Z","iopub.execute_input":"2022-12-11T10:57:05.191449Z","iopub.status.idle":"2022-12-11T10:57:11.833839Z","shell.execute_reply.started":"2022-12-11T10:57:05.191427Z","shell.execute_reply":"2022-12-11T10:57:11.832868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4 - Model","metadata":{}},{"cell_type":"code","source":"class CNN(nn.Module):\n    \n    def conv_block(self, nfilters = [16,32], ker_size = [3,3], strides = [1,2]):\n        return nn.Sequential(nn.LazyConv2d(nfilters[0], ker_size[0], strides[0]),\n                      nn.ReLU(),\n                      nn.Dropout2d(p=0.2),\n                      nn.LazyConv2d(nfilters[1], ker_size[1], strides[1]),\n                      nn.ReLU(),\n                      nn.MaxPool2d(2,2),\n                      nn.LazyBatchNorm2d())\n    \n    def linear_block(self, sizes=[64,32], dr=0.1):\n        return nn.Sequential(nn.LazyLinear(sizes[0]),\n                              nn.ReLU(),\n                              nn.Dropout(p=dr),\n                              nn.LazyLinear(sizes[1]),\n                              nn.ReLU())\n    \n    def __init__(self):\n        super().__init__()\n        self.conv1 = self.conv_block()#ker_size=[2,2])\n        self.conv2 = self.conv_block(nfilters=[32,64])\n        self.conv3 = self.conv_block(nfilters=[64,128])\n        self.embedding = nn.Flatten()\n        self.dense = self.linear_block()\n        self.classifier = nn.LazyLinear(3)\n        \n    def forward(self, x, embedding_only=False):\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.conv3(x)\n        x = self.embedding(x)\n        x = self.dense(x)\n        x = self.classifier(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.835147Z","iopub.execute_input":"2022-12-11T10:57:11.836969Z","iopub.status.idle":"2022-12-11T10:57:11.850851Z","shell.execute_reply.started":"2022-12-11T10:57:11.836929Z","shell.execute_reply":"2022-12-11T10:57:11.850058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CNN()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.852270Z","iopub.execute_input":"2022-12-11T10:57:11.852870Z","iopub.status.idle":"2022-12-11T10:57:11.871640Z","shell.execute_reply.started":"2022-12-11T10:57:11.852833Z","shell.execute_reply":"2022-12-11T10:57:11.870489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(model.cuda(), (1,HEIGHT,WIDTH))","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.873224Z","iopub.execute_input":"2022-12-11T10:57:11.873890Z","iopub.status.idle":"2022-12-11T10:57:11.897194Z","shell.execute_reply.started":"2022-12-11T10:57:11.873854Z","shell.execute_reply":"2022-12-11T10:57:11.896093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 5 - Training","metadata":{}},{"cell_type":"code","source":"optimizer = torch.optim.Adam(params=model.parameters(), lr=LEARNING_RATE)\nloss = nn.CrossEntropyLoss()\nnsteps = len(train_loader)\nacc = Accuracy(num_classes=3).cuda()\nf1 = F1Score(num_classes=3).cuda()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.900253Z","iopub.execute_input":"2022-12-11T10:57:11.900508Z","iopub.status.idle":"2022-12-11T10:57:11.910435Z","shell.execute_reply.started":"2022-12-11T10:57:11.900484Z","shell.execute_reply":"2022-12-11T10:57:11.909521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"accelerator = Accelerator(gradient_accumulation_steps=GRADIENT_ACCUMULATION_STEPS, mixed_precision=MIXED_PRECISION)\nmodel, optimizer, train_loader, val_loader = accelerator.prepare(model, optimizer, train_loader, val_loader)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.911938Z","iopub.execute_input":"2022-12-11T10:57:11.912285Z","iopub.status.idle":"2022-12-11T10:57:11.920277Z","shell.execute_reply.started":"2022-12-11T10:57:11.912252Z","shell.execute_reply":"2022-12-11T10:57:11.919231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def training_evaluation(outputs, labels, train_loss, train_acc, train_f1):\n    batchloss = loss(outputs, labels.to(int))\n    train_loss += batchloss\n    batch_acc = acc(outputs, labels.to(int))\n    train_acc += batch_acc\n    batch_f1 = f1(outputs, labels.to(int))\n    train_f1 += batch_f1\n    return (batchloss, batch_acc, batch_f1), (train_loss, train_acc, train_f1)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.921817Z","iopub.execute_input":"2022-12-11T10:57:11.922301Z","iopub.status.idle":"2022-12-11T10:57:11.930483Z","shell.execute_reply.started":"2022-12-11T10:57:11.922263Z","shell.execute_reply":"2022-12-11T10:57:11.929439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validation_evaluation(outputs, labels, val_loss, val_acc, val_f1):\n    val_loss += loss(outputs, labels.to(int))\n    val_acc += acc(outputs, labels.to(int))\n    val_f1 += f1(outputs, labels.to(int))\n    return val_loss, val_acc, val_f1","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.932375Z","iopub.execute_input":"2022-12-11T10:57:11.932706Z","iopub.status.idle":"2022-12-11T10:57:11.944357Z","shell.execute_reply.started":"2022-12-11T10:57:11.932670Z","shell.execute_reply":"2022-12-11T10:57:11.943406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def epoch_training():\n    bar = tq.tqdm(range(nsteps))\n    train_loss, train_acc, train_f1 = 0, 0, 0\n    model.train()\n    for batch in train_loader:\n        with accelerator.accumulate(model):\n            images, labels = batch\n            optimizer.zero_grad()\n            outputs = model(images)\n            batch_metrics, epoch_metrics = training_evaluation(outputs, labels, train_loss, train_acc, train_f1)\n            batchloss, batch_acc, batch_f1 = batch_metrics\n            train_loss, train_acc, train_f1 = epoch_metrics\n            accelerator.backward(batchloss)\n            optimizer.step()\n            bar.update(1)\n            logs = {\"batch loss\":batchloss.detach().item(), \"batch accuracy\":batch_acc.item(), \"batch F1\":batch_f1.item()}\n            bar.set_postfix(**logs)\n    train_loss /= len(train_loader)\n    train_acc /= len(train_loader)\n    train_f1 /= len(train_loader)\n    print(\"Training Results:\")\n    print(f\"Average Loss = {train_loss:.2f}\\nAverage Accuracy = {train_acc:.2f}\\nAverage F1Score = {train_f1:.2f}\")","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.946179Z","iopub.execute_input":"2022-12-11T10:57:11.946619Z","iopub.status.idle":"2022-12-11T10:57:11.957794Z","shell.execute_reply.started":"2022-12-11T10:57:11.946585Z","shell.execute_reply":"2022-12-11T10:57:11.956849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def epoch_validation(best_val_f1score):\n    model.eval()\n    val_loss, val_acc, val_f1 = 0, 0, 0\n    for batch in val_loader:\n        images, labels = batch\n        outputs = model(images)\n        val_loss, val_acc, val_f1 = validation_evaluation(outputs, labels, val_loss, val_acc, val_f1)\n    val_loss /= len(val_loader)\n    val_acc /= len(val_loader)\n    val_f1 /= len(val_loader)\n    print(\"Validation Results:\")\n    print(f\"Average Loss = {val_loss:.2f}\\nAverage Accuracy = {val_acc:.2f}\\nAverage F1Score = {val_f1:.2f}\")\n    if val_f1 > best_val_f1score:\n        best_val_f1score = val_f1\n        torch.save({'model_state_dict': model.state_dict()}, BEST_MODEL_PATH)\n        print(\"Validation F1Score Improved. Saving Model.\")\n    print(\"\\n\\n\")\n    return best_val_f1score","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.959377Z","iopub.execute_input":"2022-12-11T10:57:11.959805Z","iopub.status.idle":"2022-12-11T10:57:11.968639Z","shell.execute_reply.started":"2022-12-11T10:57:11.959771Z","shell.execute_reply":"2022-12-11T10:57:11.967521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train():\n    best_val_f1score = 0\n    for e in range(NEPOCHS):\n        print(f\"Epoch {e+1}:\")\n        epoch_training()\n        best_val_f1score = epoch_validation(best_val_f1score)\n    accelerator.free_memory()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.970143Z","iopub.execute_input":"2022-12-11T10:57:11.970776Z","iopub.status.idle":"2022-12-11T10:57:11.982049Z","shell.execute_reply.started":"2022-12-11T10:57:11.970740Z","shell.execute_reply":"2022-12-11T10:57:11.981095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T10:57:11.985259Z","iopub.execute_input":"2022-12-11T10:57:11.985518Z","iopub.status.idle":"2022-12-11T11:02:41.230471Z","shell.execute_reply.started":"2022-12-11T10:57:11.985494Z","shell.execute_reply":"2022-12-11T11:02:41.229398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint = torch.load(BEST_MODEL_PATH)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.238757Z","iopub.execute_input":"2022-12-11T11:02:41.239683Z","iopub.status.idle":"2022-12-11T11:02:41.264154Z","shell.execute_reply.started":"2022-12-11T11:02:41.239642Z","shell.execute_reply":"2022-12-11T11:02:41.263302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(checkpoint['model_state_dict'])","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.265452Z","iopub.execute_input":"2022-12-11T11:02:41.266075Z","iopub.status.idle":"2022-12-11T11:02:41.275212Z","shell.execute_reply.started":"2022-12-11T11:02:41.266037Z","shell.execute_reply":"2022-12-11T11:02:41.274134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 6 - XAI / Visualization","metadata":{}},{"cell_type":"markdown","source":"## 6.1 - Utilities","metadata":{}},{"cell_type":"code","source":"def get_layer_output(x, layer_name):\n    for name, layer in list(model.named_children()):\n        x = layer(x)\n        if name == layer_name:\n            return x","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.276660Z","iopub.execute_input":"2022-12-11T11:02:41.277173Z","iopub.status.idle":"2022-12-11T11:02:41.283269Z","shell.execute_reply.started":"2022-12-11T11:02:41.277074Z","shell.execute_reply":"2022-12-11T11:02:41.282384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_all_latents(name):\n    array = []\n    classes = []\n    for batch in train_loader:\n        images, labels = batch\n        output = get_layer_output(images, name).detach().cpu().numpy()\n        array.append(output)\n        labels = labels.detach().cpu().to(int).numpy()\n        classes.extend(labels)\n    for batch in val_loader:\n        images, labels = batch\n        output = get_layer_output(images, name).detach().cpu().numpy()\n        array.append(output)\n        labels = labels.detach().cpu().to(int).numpy()\n        classes.extend(labels)\n    return np.vstack(array), classes","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.284796Z","iopub.execute_input":"2022-12-11T11:02:41.285137Z","iopub.status.idle":"2022-12-11T11:02:41.294847Z","shell.execute_reply.started":"2022-12-11T11:02:41.285103Z","shell.execute_reply":"2022-12-11T11:02:41.293881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"reducer = umap.UMAP()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.296124Z","iopub.execute_input":"2022-12-11T11:02:41.296618Z","iopub.status.idle":"2022-12-11T11:02:41.303758Z","shell.execute_reply.started":"2022-12-11T11:02:41.296584Z","shell.execute_reply":"2022-12-11T11:02:41.302767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_all_reduced_latents(name):\n    latents, classes = get_all_latents(name)\n    reduced_latents = reducer.fit_transform(latents)\n    return reduced_latents, classes","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.305158Z","iopub.execute_input":"2022-12-11T11:02:41.305516Z","iopub.status.idle":"2022-12-11T11:02:41.315173Z","shell.execute_reply.started":"2022-12-11T11:02:41.305481Z","shell.execute_reply":"2022-12-11T11:02:41.314288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inverse_label_map = {ind: name for (name,ind) in LABEL_MAP.items()}","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.316426Z","iopub.execute_input":"2022-12-11T11:02:41.316880Z","iopub.status.idle":"2022-12-11T11:02:41.327906Z","shell.execute_reply.started":"2022-12-11T11:02:41.316844Z","shell.execute_reply":"2022-12-11T11:02:41.327000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize2D(reduced_latents, classes):\n    names = [inverse_label_map[i] for i in classes]\n    plt.figure(figsize=(20,8))\n    sns.scatterplot(x=reduced_latents[:,0], y=reduced_latents[:,1], hue=names, palette=\"Set1\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.329136Z","iopub.execute_input":"2022-12-11T11:02:41.329600Z","iopub.status.idle":"2022-12-11T11:02:41.338069Z","shell.execute_reply.started":"2022-12-11T11:02:41.329564Z","shell.execute_reply":"2022-12-11T11:02:41.337109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize3D(logits, classes):\n    names = [inverse_label_map[i] for i in classes]\n    fig = px.scatter_3d(x=logits[:,0], y=logits[:,1], z=logits[:,2], color=names)\n    fig.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.339434Z","iopub.execute_input":"2022-12-11T11:02:41.339864Z","iopub.status.idle":"2022-12-11T11:02:41.347935Z","shell.execute_reply.started":"2022-12-11T11:02:41.339827Z","shell.execute_reply":"2022-12-11T11:02:41.346929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def saliency(x):\n    model.eval()\n    x.requires_grad_()\n    output = model(x)\n    max_ind = output.argmax()\n    max_prob = output[0,max_ind]\n    max_prob.backward()\n    saliency = x.grad.data.abs()\n    return saliency.cpu()[0,0]","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.349185Z","iopub.execute_input":"2022-12-11T11:02:41.349653Z","iopub.status.idle":"2022-12-11T11:02:41.357666Z","shell.execute_reply.started":"2022-12-11T11:02:41.349616Z","shell.execute_reply":"2022-12-11T11:02:41.356781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_saliency(x, label):\n    s = saliency(x)\n    x = x.detach().cpu()[0,0]\n    #x = process_image(T.ToPILImage()(x.unsqueeze(dim=0)))\n    fig, axes = plt.subplots(nrows=1, ncols=3, figsize=(20,7))\n    axes[0].imshow(x, cmap=\"Blues_r\")\n    axes[1].imshow(s.cpu(), cmap=\"Reds_r\")\n    axes[2].imshow(x, cmap=\"Blues_r\")\n    axes[2].imshow(s.cpu(), cmap=\"Reds_r\", alpha=0.5)\n    axes[0].set_title(f\"Image ({label})\")\n    axes[1].set_title(f\"Saliency Map\")\n    axes[2].set_title(\"Important Regions Highlighted\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.359004Z","iopub.execute_input":"2022-12-11T11:02:41.359479Z","iopub.status.idle":"2022-12-11T11:02:41.368968Z","shell.execute_reply.started":"2022-12-11T11:02:41.359443Z","shell.execute_reply":"2022-12-11T11:02:41.368097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_saliency_maps(n):\n    for i in range(n):\n        t = random.random()\n        if t>0.5:\n            k = random.randint(0,len(traindata))\n            x, label = traindata.__getitem__(i)\n        else:\n            k = random.randint(0,len(valdata))\n            x, label = valdata.__getitem__(i)\n        x = x.unsqueeze(0).cuda()\n        label = inverse_label_map[int(label)]\n        visualize_saliency(x, label)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.370367Z","iopub.execute_input":"2022-12-11T11:02:41.370842Z","iopub.status.idle":"2022-12-11T11:02:41.379020Z","shell.execute_reply.started":"2022-12-11T11:02:41.370808Z","shell.execute_reply":"2022-12-11T11:02:41.378116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 6.2 -  Visualizing Latent Representations with UMAP","metadata":{}},{"cell_type":"code","source":"sns.set_style(\"darkgrid\")","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.380238Z","iopub.execute_input":"2022-12-11T11:02:41.380886Z","iopub.status.idle":"2022-12-11T11:02:41.388936Z","shell.execute_reply.started":"2022-12-11T11:02:41.380850Z","shell.execute_reply":"2022-12-11T11:02:41.387896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### The Convolutional Feature Extractor (Flatten Layer):","metadata":{}},{"cell_type":"code","source":"reduced_latents, classes = get_all_reduced_latents(\"embedding\")\nvisualize2D(reduced_latents, classes)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:02:41.390328Z","iopub.execute_input":"2022-12-11T11:02:41.390815Z","iopub.status.idle":"2022-12-11T11:03:25.931554Z","shell.execute_reply.started":"2022-12-11T11:02:41.390779Z","shell.execute_reply":"2022-12-11T11:03:25.930609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### The pre-Final Layer:","metadata":{}},{"cell_type":"code","source":"reduced_latents, classes = get_all_reduced_latents(\"dense\")\nvisualize2D(reduced_latents, classes)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:03:25.933078Z","iopub.execute_input":"2022-12-11T11:03:25.933753Z","iopub.status.idle":"2022-12-11T11:04:00.993690Z","shell.execute_reply.started":"2022-12-11T11:03:25.933701Z","shell.execute_reply":"2022-12-11T11:04:00.992688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 6.3 - Visualizing The Final Layer (Logits):","metadata":{}},{"cell_type":"code","source":"logits, classes = get_all_latents(\"classifier\")\n#visualize3D(logits, classes)\nnames = [inverse_label_map[i] for i in classes]\nfig = px.scatter_3d(x=logits[:,0], y=logits[:,1], z=logits[:,2], color=names)\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:04:00.995237Z","iopub.execute_input":"2022-12-11T11:04:00.995700Z","iopub.status.idle":"2022-12-11T11:04:28.247264Z","shell.execute_reply.started":"2022-12-11T11:04:00.995654Z","shell.execute_reply":"2022-12-11T11:04:28.246185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 6.4 - Saliency Maps","metadata":{}},{"cell_type":"code","source":"sns.set_style(\"white\")","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:04:28.249747Z","iopub.execute_input":"2022-12-11T11:04:28.250328Z","iopub.status.idle":"2022-12-11T11:04:28.255143Z","shell.execute_reply.started":"2022-12-11T11:04:28.250290Z","shell.execute_reply":"2022-12-11T11:04:28.254174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_saliency_maps(10)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:04:28.258295Z","iopub.execute_input":"2022-12-11T11:04:28.258571Z","iopub.status.idle":"2022-12-11T11:04:35.434695Z","shell.execute_reply.started":"2022-12-11T11:04:28.258546Z","shell.execute_reply":"2022-12-11T11:04:35.433778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 7 - Tabular ML Approach with Manually Chosen Features","metadata":{}},{"cell_type":"markdown","source":"### The features used are:\n* Mean of pixel values\n* Std of pixel values\n* Mean of pixel values after equalization\n* Std of pixel values after equalization","metadata":{}},{"cell_type":"code","source":"def get_features(x):\n    mean = x.mean().item()\n    std = x.std().item()\n    x = process_image((x*255).type(torch.uint8))/255\n    mean_eq = x.mean().item()\n    std_eq = x.std().item()\n    return mean, std, mean_eq, std_eq","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:04:35.436260Z","iopub.execute_input":"2022-12-11T11:04:35.436822Z","iopub.status.idle":"2022-12-11T11:04:35.444433Z","shell.execute_reply.started":"2022-12-11T11:04:35.436787Z","shell.execute_reply":"2022-12-11T11:04:35.443350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_tabular_dataset(traindata):\n    tabdata = np.zeros(shape=(len(traindata),5))\n    for i in range(len(traindata)):\n        x = traindata.__getitem__(i)\n        tabdata[i,:4] = get_features(x[0])\n        tabdata[i,4] = int(x[1].item())\n    return pd.DataFrame(tabdata, columns=[\"mean\",\"std\",\"mean_equalized\",\"std_equalized\",\"label\"])","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:04:35.445963Z","iopub.execute_input":"2022-12-11T11:04:35.446645Z","iopub.status.idle":"2022-12-11T11:04:35.456104Z","shell.execute_reply.started":"2022-12-11T11:04:35.446605Z","shell.execute_reply":"2022-12-11T11:04:35.455046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tabdata = create_tabular_dataset(traindata)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:04:35.457727Z","iopub.execute_input":"2022-12-11T11:04:35.458365Z","iopub.status.idle":"2022-12-11T11:05:10.113756Z","shell.execute_reply.started":"2022-12-11T11:04:35.458325Z","shell.execute_reply":"2022-12-11T11:05:10.112769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tabdata.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:05:10.115078Z","iopub.execute_input":"2022-12-11T11:05:10.115438Z","iopub.status.idle":"2022-12-11T11:05:10.128038Z","shell.execute_reply.started":"2022-12-11T11:05:10.115402Z","shell.execute_reply":"2022-12-11T11:05:10.126971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x, y = tabdata.drop(\"label\", axis=1), tabdata[\"label\"]\nxtrain, xtest, ytrain, ytest = train_test_split(x, y, test_size=0.1)","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:05:10.129759Z","iopub.execute_input":"2022-12-11T11:05:10.130507Z","iopub.status.idle":"2022-12-11T11:05:10.142158Z","shell.execute_reply.started":"2022-12-11T11:05:10.130440Z","shell.execute_reply":"2022-12-11T11:05:10.141015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rf = RandomForestClassifier().fit(xtrain,ytrain)\nypred_tr = rf.predict(xtrain)\nypred_ts = rf.predict(xtest)\nprint(f\"Training Results:\\n{classification_report(ytrain,ypred_tr)}\\n\\nTesting Results:\\n{classification_report(ytest,ypred_ts)}\")","metadata":{"execution":{"iopub.status.busy":"2022-12-11T11:05:10.144063Z","iopub.execute_input":"2022-12-11T11:05:10.144577Z","iopub.status.idle":"2022-12-11T11:05:10.603447Z","shell.execute_reply.started":"2022-12-11T11:05:10.144540Z","shell.execute_reply":"2022-12-11T11:05:10.602249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}