{"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":"!pip install self-supervised -Uq","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:45.653360Z","iopub.execute_input":"2021-06-23T06:50:45.653879Z","iopub.status.idle":"2021-06-23T06:50:51.626291Z","shell.execute_reply.started":"2021-06-23T06:50:45.653805Z","shell.execute_reply":"2021-06-23T06:50:51.625232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nfrom fastai import *\nfrom fastai.data.all import *\nfrom fastai.vision.all import *\nfrom self_supervised.augmentations import *\nfrom self_supervised.layers import *\nfrom self_supervised.vision.swav import *\nfrom torch.utils.data import Dataset\nimport pandas as pd\nimport torch\nimport torchvision.models as models\nimport torchvision.transforms as transforms","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-06-23T06:50:51.633601Z","iopub.execute_input":"2021-06-23T06:50:51.633933Z","iopub.status.idle":"2021-06-23T06:50:52.891457Z","shell.execute_reply.started":"2021-06-23T06:50:51.633895Z","shell.execute_reply":"2021-06-23T06:50:52.890529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_x(x): return data_dir +'/train/'+ x['id']+'.tif'\n\ndef get_dls(size, bs, df):\n    \n    db = DataBlock(blocks = (ImageBlock(), CategoryBlock()),\n              get_x = get_x, get_y=ColReader('label'),\n              splitter=ColSplitter())\n    \n    dls = db.dataloaders(df, bs=bs)\n    return dls","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:52.893256Z","iopub.execute_input":"2021-06-23T06:50:52.893585Z","iopub.status.idle":"2021-06-23T06:50:52.902157Z","shell.execute_reply.started":"2021-06-23T06:50:52.893554Z","shell.execute_reply":"2021-06-23T06:50:52.901240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:52.903480Z","iopub.execute_input":"2021-06-23T06:50:52.903993Z","iopub.status.idle":"2021-06-23T06:50:52.911435Z","shell.execute_reply.started":"2021-06-23T06:50:52.903956Z","shell.execute_reply":"2021-06-23T06:50:52.910351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = \"../input/histopathologic-cancer-detection\"","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:52.912811Z","iopub.execute_input":"2021-06-23T06:50:52.913237Z","iopub.status.idle":"2021-06-23T06:50:52.919715Z","shell.execute_reply.started":"2021-06-23T06:50:52.913203Z","shell.execute_reply":"2021-06-23T06:50:52.918576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('../input/histopathologic-cancer-detection/train_labels.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:52.921162Z","iopub.execute_input":"2021-06-23T06:50:52.921594Z","iopub.status.idle":"2021-06-23T06:50:53.108327Z","shell.execute_reply.started":"2021-06-23T06:50:52.921559Z","shell.execute_reply":"2021-06-23T06:50:53.107305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GaussianNoise:\n    \"\"\"Applies random Gaussian noise to a tensor.\n\n    The intensity of the noise is dependent on the mean of the pixel values.\n    See https://arxiv.org/pdf/2101.04909.pdf for more information.\n\n    \"\"\"\n\n    def __call__(self, sample: torch.Tensor) -> torch.Tensor:\n        mu = sample.mean()\n        snr = np.random.randint(low=4, high=8)\n        sigma = mu / snr\n        noise = torch.normal(torch.zeros(sample.shape), sigma)\n        return sample + noise","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:53.109774Z","iopub.execute_input":"2021-06-23T06:50:53.110138Z","iopub.status.idle":"2021-06-23T06:50:53.115715Z","shell.execute_reply.started":"2021-06-23T06:50:53.110100Z","shell.execute_reply":"2021-06-23T06:50:53.114746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transformer = transforms.Compose([\n    transforms.Grayscale(num_output_channels=3),\n    transforms.RandomResizedCrop(size=96, scale=(0.2, 1.0)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.GaussianBlur(11),\n    transforms.ToTensor(),\n    GaussianNoise(),\n])","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:53.119011Z","iopub.execute_input":"2021-06-23T06:50:53.119425Z","iopub.status.idle":"2021-06-23T06:50:53.127203Z","shell.execute_reply.started":"2021-06-23T06:50:53.119385Z","shell.execute_reply":"2021-06-23T06:50:53.126290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example_image_name = 'f38a6374c348f90b587e046aac6079959adf3835.tif'\nexample_image_path = os.path.join(data_dir, \"train/\"+example_image_name)\nexample_image = Image.open(example_image_path)\n\n# torch transform returns a 3 x W x H image, we only show one color channel\naugmented_image_1 = data_transformer(example_image).numpy()[0]\naugmented_image_2 = data_transformer(example_image).numpy()[0]\n\nfig, axs = plt.subplots(1, 3)\n\naxs[0].imshow(example_image)\naxs[0].set_axis_off()\naxs[0].set_title('Original Image')\n\naxs[1].imshow(augmented_image_1)\naxs[1].set_axis_off()\naxs[1].set_title('Augmented-1')\n\naxs[2].imshow(augmented_image_2)\naxs[2].set_axis_off()\naxs[2].set_title('Augmented-2')","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:53.129465Z","iopub.execute_input":"2021-06-23T06:50:53.129961Z","iopub.status.idle":"2021-06-23T06:50:53.429237Z","shell.execute_reply.started":"2021-06-23T06:50:53.129924Z","shell.execute_reply":"2021-06-23T06:50:53.428383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32\n\ndir_of_files = data_dir+\"/train\" \nfilenames = os.listdir(dir_of_files) \nfiles = [f.replace(\".tif\",\"\") for f in filenames]\n\ncut = int(0.8 * len(files))\n\ntrain_files = files[:cut] \nvalid_files = files[cut:]\n\n# For feature extration using 20% of train data and 10% of validation data\nfe_train_len = int(0.2*len(train_files))\nfe_valid_len = int(0.1*len(valid_files))\n\nfe_train_files = train_files[:fe_train_len]\nfe_valid_files = valid_files[:fe_valid_len]\n\nprint(len(fe_train_files))\nprint(len(fe_valid_files))","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:53.430694Z","iopub.execute_input":"2021-06-23T06:50:53.431061Z","iopub.status.idle":"2021-06-23T06:50:53.666347Z","shell.execute_reply.started":"2021-06-23T06:50:53.431023Z","shell.execute_reply":"2021-06-23T06:50:53.665400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['is_valid'] = False\ndf['is_fe'] = False\ndf.loc[df['id'].isin(fe_valid_files), 'is_valid'] = True\ndf.loc[df['id'].isin(fe_valid_files), 'is_fe'] = True\ndf.loc[df['id'].isin(fe_train_files), 'is_fe'] = True\n\ndf.groupby('is_valid').label.value_counts()\ndf.groupby('is_fe').label.value_counts()","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:53.667664Z","iopub.execute_input":"2021-06-23T06:50:53.668194Z","iopub.status.idle":"2021-06-23T06:50:53.795238Z","shell.execute_reply.started":"2021-06-23T06:50:53.668154Z","shell.execute_reply":"2021-06-23T06:50:53.794425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"size=96","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:53.797054Z","iopub.execute_input":"2021-06-23T06:50:53.797297Z","iopub.status.idle":"2021-06-23T06:50:53.803412Z","shell.execute_reply.started":"2021-06-23T06:50:53.797273Z","shell.execute_reply":"2021-06-23T06:50:53.802556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision import models\n\nclass BasicBlock(nn.Module):\n    expansion = 1\n\n    def __init__(self, inplanes, planes, stride=1, downsample=None):\n        super().__init__()\n        self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=3, stride=stride,\n                     padding=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=1,\n                     padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(planes)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        identity = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        if self.downsample is not None:\n            identity = self.downsample(x)\n\n        out += identity\n        out = self.relu(out)\n\n        return out\n        \nclass ResNet(nn.Module):\n\n    def __init__(self, block, layers, num_classes=1000):\n        super().__init__()\n        \n        self.inplanes = 64\n\n        self.conv1 = nn.Conv2d(3, self.inplanes, kernel_size=7, stride=2, padding=3,\n                               bias=False)\n        self.bn1 = nn.BatchNorm2d(self.inplanes)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        \n        self.layer1 = self._make_layer(block, 64, layers[0])\n        self.layer2 = self._make_layer(block, 128, layers[1], stride=2)\n        self.layer3 = self._make_layer(block, 256, layers[2], stride=2)\n        self.layer4 = self._make_layer(block, 512, layers[3], stride=2)\n        \n        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n        self.fc = nn.Linear(512 , num_classes)\n\n\n    def _make_layer(self, block, planes, blocks, stride=1):\n        downsample = None  \n   \n        if stride != 1 or self.inplanes != planes:\n            downsample = nn.Sequential(\n                nn.Conv2d(self.inplanes, planes, 1, stride, bias=False),\n                nn.BatchNorm2d(planes),\n            )\n\n        layers = []\n        layers.append(block(self.inplanes, planes, stride, downsample))\n        \n        self.inplanes = planes\n        \n        for _ in range(1, blocks):\n            layers.append(block(self.inplanes, planes))\n\n        return nn.Sequential(*layers)\n    \n    \n    def forward(self, x):\n        x = self.conv1(x)           # 224x224\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)         # 112x112\n\n        x = self.layer1(x)          # 56x56\n        x = self.layer2(x)          # 28x28\n        x = self.layer3(x)          # 14x14\n        x = self.layer4(x)          # 7x7\n\n        x = self.avgpool(x)         # 1x1\n        x = torch.flatten(x, 1)     # remove 1 X 1 grid and make vector of tensor shape \n        x = self.fc(x)\n\n        return x\n\ndef resnet_temp():\n    layers=[1, 1, 1, 1]\n    model = ResNet(BasicBlock, layers)\n    return model\n\ndef weights_copy(custom_model, resnet18):\n    # print(custom_model.state_dict)\n    model_custom.conv1 = resnet18.conv1\n    model_custom.bn1 = resnet18.bn1\n    model_custom.maxpool = resnet18.maxpool\n  \n    model_custom.layer1[0] = resnet18.layer1[0]\n    model_custom.layer2[0] = resnet18.layer2[0]\n    model_custom.layer3[0] = resnet18.layer3[0]\n    model_custom.layer4[0] = resnet18.layer4[0]\n\n    model_custom.avgpool = resnet18.avgpool\n    model_custom.fc = resnet18.fc\n\n    return model_custom\n\nmodel_custom = resnet_temp()\nprint(model_custom)\n","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:53.804836Z","iopub.execute_input":"2021-06-23T06:50:53.805192Z","iopub.status.idle":"2021-06-23T06:50:53.871984Z","shell.execute_reply.started":"2021-06-23T06:50:53.805159Z","shell.execute_reply":"2021-06-23T06:50:53.871111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"arch = \"xresnet18\"\nencoder = models.resnet18(pretrained=True)\nencoder = weights_copy(model_custom, encoder)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:53.873323Z","iopub.execute_input":"2021-06-23T06:50:53.873712Z","iopub.status.idle":"2021-06-23T06:50:54.182370Z","shell.execute_reply.started":"2021-06-23T06:50:53.873675Z","shell.execute_reply":"2021-06-23T06:50:54.181531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_for_fe = df[df['is_fe'] == True]\nlen(df_for_fe)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:54.183853Z","iopub.execute_input":"2021-06-23T06:50:54.184211Z","iopub.status.idle":"2021-06-23T06:50:54.197341Z","shell.execute_reply.started":"2021-06-23T06:50:54.184174Z","shell.execute_reply":"2021-06-23T06:50:54.196178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = get_dls(size, batch_size, df_for_fe)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:50:54.198606Z","iopub.execute_input":"2021-06-23T06:50:54.198946Z","iopub.status.idle":"2021-06-23T06:51:00.308354Z","shell.execute_reply.started":"2021-06-23T06:50:54.198910Z","shell.execute_reply":"2021-06-23T06:51:00.307500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = create_swav_model(encoder)\naug_pipelines = get_swav_aug_pipelines(num_crops=[2,6],\n                                       crop_sizes=[size,int(3/4*size)], \n                                       min_scales=[0.25,0.2],\n                                       max_scales=[1.0,0.35],\n                                       rotate=False, jitter=False, bw=False, blur=False) ","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:51:00.309651Z","iopub.execute_input":"2021-06-23T06:51:00.310178Z","iopub.status.idle":"2021-06-23T06:51:00.396618Z","shell.execute_reply.started":"2021-06-23T06:51:00.310137Z","shell.execute_reply":"2021-06-23T06:51:00.395770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"K = batch_size*2**4\ncbs=[SWAV(aug_pipelines, crop_assgn_ids=[0,1], K=K, queue_start_pct=0.5, temp=0.1)]","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:51:00.397874Z","iopub.execute_input":"2021-06-23T06:51:00.398368Z","iopub.status.idle":"2021-06-23T06:51:00.403689Z","shell.execute_reply.started":"2021-06-23T06:51:00.398327Z","shell.execute_reply":"2021-06-23T06:51:00.402594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn = Learner(dls, model, cbs=cbs)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:51:00.404953Z","iopub.execute_input":"2021-06-23T06:51:00.405363Z","iopub.status.idle":"2021-06-23T06:51:00.413540Z","shell.execute_reply.started":"2021-06-23T06:51:00.405325Z","shell.execute_reply":"2021-06-23T06:51:00.412721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"b = dls.one_batch()\nlearn._split(b)\nlearn('before_batch')\nlearn.swav.show(n=5);","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:51:00.414772Z","iopub.execute_input":"2021-06-23T06:51:00.415119Z","iopub.status.idle":"2021-06-23T06:51:02.989425Z","shell.execute_reply.started":"2021-06-23T06:51:00.415081Z","shell.execute_reply":"2021-06-23T06:51:02.988647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr, wd =1e-4, 1e-2\nepochs =10\nlearn.unfreeze()\nlearn.fit_flat_cos(epochs, lr, wd, pct_start=0.5, cbs=EarlyStoppingCallback(monitor='train_loss', min_delta=0.1, patience=2))","metadata":{"execution":{"iopub.status.busy":"2021-06-23T06:51:02.990668Z","iopub.execute_input":"2021-06-23T06:51:02.991132Z","iopub.status.idle":"2021-06-23T07:01:52.266146Z","shell.execute_reply.started":"2021-06-23T06:51:02.991093Z","shell.execute_reply":"2021-06-23T07:01:52.265103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output_path = \"./models/\"\nsave_name = f'swav_{size}_epc{epochs}'\nlearn.save(save_name)\ntorch.save(learn.model.encoder.state_dict(), output_path+save_name+'_encoder.pth')\nlearn.recorder.plot_loss()","metadata":{"execution":{"iopub.status.busy":"2021-06-23T07:01:52.267744Z","iopub.execute_input":"2021-06-23T07:01:52.268118Z","iopub.status.idle":"2021-06-23T07:01:52.718392Z","shell.execute_reply.started":"2021-06-23T07:01:52.268078Z","shell.execute_reply":"2021-06-23T07:01:52.717415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Evaluating**","metadata":{}},{"cell_type":"code","source":"import albumentations as A \nfrom albumentations.pytorch.transforms import ToTensorV2\n\nimport torch\nimport torchvision.transforms as transforms\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom tqdm import tqdm\nimport cv2\nimport gc","metadata":{"execution":{"iopub.status.busy":"2021-06-23T07:02:58.081335Z","iopub.execute_input":"2021-06-23T07:02:58.081696Z","iopub.status.idle":"2021-06-23T07:02:58.087614Z","shell.execute_reply.started":"2021-06-23T07:02:58.081666Z","shell.execute_reply":"2021-06-23T07:02:58.086599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_size = 96\nbatch_size = 24","metadata":{"execution":{"iopub.status.busy":"2021-06-23T07:02:58.186259Z","iopub.execute_input":"2021-06-23T07:02:58.186571Z","iopub.status.idle":"2021-06-23T07:02:58.190504Z","shell.execute_reply.started":"2021-06-23T07:02:58.186541Z","shell.execute_reply":"2021-06-23T07:02:58.189288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transformer = transforms.Compose([transforms.ToTensor()])","metadata":{"execution":{"iopub.status.busy":"2021-06-23T07:02:58.336096Z","iopub.execute_input":"2021-06-23T07:02:58.336420Z","iopub.status.idle":"2021-06-23T07:02:58.340949Z","shell.execute_reply.started":"2021-06-23T07:02:58.336390Z","shell.execute_reply":"2021-06-23T07:02:58.339677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 24\n\ndir_of_files = data_dir+\"/train\" \nfilenames = os.listdir(dir_of_files) \nfiles = [f.replace(\".tif\",\"\") for f in filenames]\n\ncut = int(0.8 * len(files))\n\ntrain_files = files[:cut] \nvalid_files = files[cut:]\n\n# For downstreaming using 10% of train data and 100% of validation data\ndownstream_train_len = int(0.1*len(train_files))\ndownstream_valid_len = int(1*len(valid_files))\n\ndownstream_train_files = train_files[:downstream_train_len]\ndownstream_valid_files = valid_files[:downstream_valid_len]\n\nprint(len(downstream_train_files))\nprint(len(downstream_valid_files))","metadata":{"execution":{"iopub.status.busy":"2021-06-23T09:42:30.384482Z","iopub.execute_input":"2021-06-23T09:42:30.384912Z","iopub.status.idle":"2021-06-23T09:42:30.651349Z","shell.execute_reply.started":"2021-06-23T09:42:30.384877Z","shell.execute_reply":"2021-06-23T09:42:30.650352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['is_valid'] = False\ndf['is_downstream'] = False\ndf.loc[df['id'].isin(downstream_valid_files), 'is_valid'] = True\ndf.loc[df['id'].isin(downstream_valid_files), 'is_downstream'] = True\ndf.loc[df['id'].isin(downstream_train_files), 'is_downstream'] = True\n\ndf.groupby('is_valid').label.value_counts()\ndf.groupby('is_downstream').label.value_counts()","metadata":{"execution":{"iopub.status.busy":"2021-06-23T09:42:30.652883Z","iopub.execute_input":"2021-06-23T09:42:30.653248Z","iopub.status.idle":"2021-06-23T09:42:30.819142Z","shell.execute_reply.started":"2021-06-23T09:42:30.653213Z","shell.execute_reply":"2021-06-23T09:42:30.818139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_for_downstream = df[df['is_downstream'] == True]\nlen(df_for_downstream)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T09:42:30.821232Z","iopub.execute_input":"2021-06-23T09:42:30.821647Z","iopub.status.idle":"2021-06-23T09:42:30.845372Z","shell.execute_reply.started":"2021-06-23T09:42:30.821605Z","shell.execute_reply":"2021-06-23T09:42:30.844515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = get_dls(size, batch_size, df_for_downstream)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T09:42:30.846984Z","iopub.execute_input":"2021-06-23T09:42:30.847589Z","iopub.status.idle":"2021-06-23T09:42:33.211179Z","shell.execute_reply.started":"2021-06-23T09:42:30.847551Z","shell.execute_reply":"2021-06-23T09:42:33.210335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optdict = dict(sqr_mom=0.99,mom=0.95,beta=0.,eps=1e-4)\nopt_func = partial(ranger, **optdict)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T09:42:33.214183Z","iopub.execute_input":"2021-06-23T09:42:33.214433Z","iopub.status.idle":"2021-06-23T09:42:33.220323Z","shell.execute_reply.started":"2021-06-23T09:42:33.214409Z","shell.execute_reply":"2021-06-23T09:42:33.219544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def split_func(m): return L(m[0], m[1]).map(params)\n\ndef create_learner(size=96, arch='resnet50', encoder_path=\"./models/swav_96_epc10_encoder.pth\"):\n    \n    pretrained_encoder = torch.load(encoder_path)\n    #encoder = create_encoder(arch, pretrained=False, n_in=3)\n    encoder.load_state_dict(pretrained_encoder)\n    nf = encoder(torch.randn(2,3,224,224, device=device)).size(-1)\n    classifier = create_cls_module(nf, dls.c)\n    model = nn.Sequential(encoder, classifier)\n    learn = Learner(dls, model, opt_func=opt_func, splitter=split_func,\n                metrics=[accuracy], loss_func=CrossEntropyLossFlat())\n    return learn","metadata":{"execution":{"iopub.status.busy":"2021-06-23T09:42:33.221664Z","iopub.execute_input":"2021-06-23T09:42:33.222056Z","iopub.status.idle":"2021-06-23T09:42:33.247571Z","shell.execute_reply.started":"2021-06-23T09:42:33.222022Z","shell.execute_reply":"2021-06-23T09:42:33.246314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def finetune(size, epochs, arch, encoder_path, lr=1e-2, wd=1e-2):\n    learn = create_learner(size, arch, encoder_path)\n    learn.unfreeze()\n    learn.fit_flat_cos(epochs, lr, wd=wd)\n    final_acc = learn.recorder.values[-1][-2]\n    return final_acc","metadata":{"execution":{"iopub.status.busy":"2021-06-23T09:42:33.249525Z","iopub.execute_input":"2021-06-23T09:42:33.249916Z","iopub.status.idle":"2021-06-23T09:42:33.258005Z","shell.execute_reply.started":"2021-06-23T09:42:33.249879Z","shell.execute_reply":"2021-06-23T09:42:33.257210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc = []\nruns = 1\nfor i in range(runs): acc += [finetune(96, epochs=50, arch='resnet50', encoder_path='./models/swav_96_epc10_encoder.pth')]","metadata":{"execution":{"iopub.status.busy":"2021-06-23T09:42:33.259903Z","iopub.execute_input":"2021-06-23T09:42:33.260222Z","iopub.status.idle":"2021-06-23T11:22:29.096824Z","shell.execute_reply.started":"2021-06-23T09:42:33.260197Z","shell.execute_reply":"2021-06-23T11:22:29.095710Z"},"trusted":true},"execution_count":null,"outputs":[]}]}