{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nimport torch\nimport cv2","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"label_keys = [\n    'ETT - Abnormal',\n    'ETT - Borderline',\n    'ETT - Normal',\n    'NGT - Abnormal',\n    'NGT - Borderline',\n    'NGT - Incompletely Imaged',\n    'NGT - Normal',\n    'CVC - Abnormal',\n    'CVC - Borderline',\n    'CVC - Normal',\n    'Swan Ganz Catheter Present'\n]\n\ndef get_image_from_id(uid, shape=None):\n    img = mpimg.imread(\"../input/ranzcr-clip-catheter-line-classification/train/\" + uid + \".jpg\")\n    if shape is None: return img\n    return cv2.resize(img, dsize=shape, interpolation=cv2.INTER_CUBIC)\n\ndef get_image(sample, shape=None):\n    return get_image_from_id(sample[\"StudyInstanceUID\"], shape)\n\ndef get_labels(sample):\n    return [sample[key] for key in label_keys]\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class Dataset(torch.utils.data.Dataset):\n    'Characterizes a dataset for PyTorch'\n    def __init__(self, df):\n        list_IDs = []\n        labels = {}\n        for _, row in df.iterrows():\n            ID = row['StudyInstanceUID']\n            list_IDs.append(ID)\n            labels[ID] = get_labels(row)\n            \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        X = torch.from_numpy(get_image_from_id(ID, (128, 128))[np.newaxis, :, :]).float()\n        y = torch.from_numpy(np.array(self.labels[ID])).float()\n        return X, y","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", device)\n\ndf = pd.read_csv(\"../input/ranzcr-clip-catheter-line-classification/train.csv\")\n\nmsk = np.random.rand(len(df)) < 0.8\npartition = {}\npartition['train'] = df[msk]\npartition['validation'] = df[~msk]\n\n# Parameters\nparams = {\n    'batch_size': 64,\n    'shuffle': True,\n    'num_workers': 6\n}\n\n# Generators\ntraining_set = Dataset(partition['train'])\ntraining_generator = torch.utils.data.DataLoader(training_set, **params)\n\nvalidation_set = Dataset(partition['validation'])\nvalidation_generator = torch.utils.data.DataLoader(validation_set, **params)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\n\nimport torch.optim as optim\n\nnum_outputs = len(label_keys)\nview_size = 16*29*29\n\nclass Net(nn.Module):\n    def __init__(self):\n        super(Net, self).__init__()\n        self.conv1 = nn.Conv2d(1, 6, 5)\n        self.pool = nn.MaxPool2d(2, 2)\n        self.conv2 = nn.Conv2d(6, 16, 5)\n        self.fc1 = nn.Linear(view_size, 120)\n        self.fc2 = nn.Linear(120, 84)\n        self.fc3 = nn.Linear(84, num_outputs)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = F.relu(x)\n        x = self.pool(x)\n        \n        x = self.conv2(x)\n        x = F.relu(x)\n        \n        x = self.pool(x)\n        x = x.view(-1, view_size)\n        \n        x = self.fc1(x)\n        x = F.relu(x)\n        \n        x = self.fc2(x)\n        x = F.relu(x)\n\n        x = self.fc3(x)\n        return torch.round(x)\n        \nnet = Net()\nnet = net.to(device)\n\ncriterion = nn.MSELoss()\noptimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)  \n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"max_epochs = 5\nnum_generated = len(training_generator)\n\nfor epoch in range(max_epochs):  # loop over the dataset multiple times\n    running_loss = 0.0\n    for i, data in enumerate(training_generator, 0):\n        # get the inputs; data is a list of [inputs, labels]\n        inputs, labels = data[0].to(device), data[1].to(device)\n        # zero the parameter gradients\n        optimizer.zero_grad()\n\n        # forward + backward + optimize\n        outputs = net(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        # print statistics\n        running_loss += loss.item()\n        progress = (i + epoch * num_generated) / num_generated\n        loss_str = \"Epoch: \" + str(epoch) + \" - Progress: \" + str(round(progress*1000)/10) + \"%\"\n        loss_str += \" - Loss: \" + str(round(running_loss*1000))\n\n        print(loss_str)\n        running_loss = 0.0\n\nprint('Finished Training')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"do_skip = False\nfor local_batch, local_labels in validation_generator:\n    if do_skip:\n        break\n    else:\n        do_skip = True\n            \n    inputs, labels = local_batch.to(device), local_labels.to(device)\n    optimizer.zero_grad()\n    outputs = net(inputs)\n    print(\"Labels:  \", labels)\n    print(\"Outputs: \", outputs)\n    loss = criterion(outputs, labels)\n    loss.backward()\n    optimizer.step()\n\n    # print statistics\n    running_loss += loss.item()\n\nprint(\"Loss: \", running_loss)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}