{"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 torchvision==0.10.0","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-10-19T14:53:06.422545Z","iopub.execute_input":"2021-10-19T14:53:06.423170Z","iopub.status.idle":"2021-10-19T14:54:22.427095Z","shell.execute_reply.started":"2021-10-19T14:53:06.423035Z","shell.execute_reply":"2021-10-19T14:54:22.426087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torchvision","metadata":{"execution":{"iopub.status.busy":"2021-10-19T14:54:42.655686Z","iopub.execute_input":"2021-10-19T14:54:42.656095Z","iopub.status.idle":"2021-10-19T14:54:43.814400Z","shell.execute_reply.started":"2021-10-19T14:54:42.656055Z","shell.execute_reply":"2021-10-19T14:54:43.813269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(torch.__version__,\ntorchvision.__version__)","metadata":{"execution":{"iopub.status.busy":"2021-10-19T14:54:46.879194Z","iopub.execute_input":"2021-10-19T14:54:46.879575Z","iopub.status.idle":"2021-10-19T14:54:46.886618Z","shell.execute_reply.started":"2021-10-19T14:54:46.879537Z","shell.execute_reply":"2021-10-19T14:54:46.885277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Clone Recursion Pharma's utilities\n!git clone https://github.com/recursionpharma/rxrx1-utils.git\n\nsys.path.append('/kaggle/working/rxrx1-utils')\nimport rxrx.io as rio","metadata":{"execution":{"iopub.status.busy":"2021-10-19T14:54:49.958196Z","iopub.execute_input":"2021-10-19T14:54:49.959187Z","iopub.status.idle":"2021-10-19T14:54:58.495838Z","shell.execute_reply.started":"2021-10-19T14:54:49.959131Z","shell.execute_reply":"2021-10-19T14:54:58.494611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Loading in and resizing the images to 224x224 for resnet\n\ndef load_and_resize(dataset, experiment, plate, well, site):\n    img = rio.load_site(dataset, experiment, plate, well, site, base_path='../input/recursion-cellular-image-classification/')\n    resized = cv2.resize(img, (224,224)).astype(np.float32)\n    resized = torch.from_numpy(resized).permute(2,0,1)\n    return resized","metadata":{"execution":{"iopub.status.busy":"2021-10-19T14:55:01.215483Z","iopub.execute_input":"2021-10-19T14:55:01.215869Z","iopub.status.idle":"2021-10-19T14:55:01.222636Z","shell.execute_reply.started":"2021-10-19T14:55:01.215823Z","shell.execute_reply":"2021-10-19T14:55:01.221912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Had to make my own dataset class, the torch example using ImageFolder\n#and torch transforms wouldn't work on my dataset\n\nclass Dataset(torch.utils.data.Dataset):\n    def __init__(self, list_IDs, labels):\n        'Initialization'\n        self.labels = labels\n        self.list_IDs = list_IDs\n\n    def __len__(self):\n        'Denotes the total number of samples'\n        return len(self.list_IDs)\n\n    def __getitem__(self, index):\n        'Generates one sample of data'\n        # Select sample\n        ID = self.list_IDs[index]\n\n        # Load data and get label\n        exp = train_df.experiment.iloc[ID]\n        plate = train_df.plate.iloc[ID]\n        well = train_df.well.iloc[ID]\n        site = train_df.site.iloc[ID]\n        \n        X = load_and_resize('train', exp, plate, well, site)\n        y = self.labels[ID]\n\n        return X, y","metadata":{"execution":{"iopub.status.busy":"2021-10-19T15:17:46.329058Z","iopub.execute_input":"2021-10-19T15:17:46.330116Z","iopub.status.idle":"2021-10-19T15:17:46.341299Z","shell.execute_reply.started":"2021-10-19T15:17:46.330059Z","shell.execute_reply":"2021-10-19T15:17:46.339702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Getting a subset of the data to test with\n\ntrain_df = pd.read_csv('../input/recursion-cellular-image-classification/train.csv')\ntrain_df = train_df.iloc[0:4]\ntrain_df['site'] = 1\nindexes = [0,1,2,3]\nsirnas = [250, 60, 43, 20]","metadata":{"execution":{"iopub.status.busy":"2021-10-19T15:17:49.219847Z","iopub.execute_input":"2021-10-19T15:17:49.221114Z","iopub.status.idle":"2021-10-19T15:17:49.288306Z","shell.execute_reply.started":"2021-10-19T15:17:49.221051Z","shell.execute_reply":"2021-10-19T15:17:49.286946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Importing a pre-trained ResNet model for transfer learning and freezing all but the last layer\nresnet = torchvision.models.resnet18(pretrained=True)\nfor param in resnet.parameters():\n    param.requires_grad = False\n    \n#Redefining the final fully connected layer\ninfeat = resnet.fc.in_features\nnclasses = 1139\nresnet.fc = nn.Linear(infeat, nclasses)\n\n#Adding a convolutional layer to match the channels to the resnet model\nfirst_conv_layer = [torch.nn.Conv2d(6, 3, kernel_size=3, stride=1, padding=1, dilation=1, groups=1, bias=True)]\nfirst_conv_layer.extend(list(resnet.children()))\nresnet = torch.nn.Sequential(*first_conv_layer)\n\nresnet = resnet.to(device)\n\ncriterion = nn.CrossEntropyLoss()\n\noptimizer_conv = torch.optim.SGD(filter(lambda p: p.requires_grad, resnet.parameters()), lr=0.001, momentum=0.9)\n\n# Decay LR by a factor of 0.1 every 7 epochs\nexp_lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer_conv, step_size=7, gamma=0.1)","metadata":{"execution":{"iopub.status.busy":"2021-10-19T15:17:51.879526Z","iopub.execute_input":"2021-10-19T15:17:51.880129Z","iopub.status.idle":"2021-10-19T15:17:52.251383Z","shell.execute_reply.started":"2021-10-19T15:17:51.880090Z","shell.execute_reply":"2021-10-19T15:17:52.250681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloader = torch.utils.data.DataLoader(Dataset(indexes, sirnas), batch_size=4, shuffle=True, num_workers=4)","metadata":{"execution":{"iopub.status.busy":"2021-10-19T15:17:57.086027Z","iopub.execute_input":"2021-10-19T15:17:57.086367Z","iopub.status.idle":"2021-10-19T15:17:57.092582Z","shell.execute_reply.started":"2021-10-19T15:17:57.086329Z","shell.execute_reply":"2021-10-19T15:17:57.091070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\nfor inputs, labels in dataloader:\n    inputs = inputs.to(device)\n    labels = torch.tensor(labels).to(device)\n    outputs = resnet(inputs)","metadata":{"execution":{"iopub.status.busy":"2021-10-19T15:18:00.387292Z","iopub.execute_input":"2021-10-19T15:18:00.387648Z","iopub.status.idle":"2021-10-19T15:18:01.370857Z","shell.execute_reply.started":"2021-10-19T15:18:00.387613Z","shell.execute_reply":"2021-10-19T15:18:01.368380Z"},"trusted":true},"execution_count":null,"outputs":[]}]}