{"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":"# Preparation\nPlease add data https://www.kaggle.com/datasets/jirkaborovec/stroke-blood-clot-origin-1k-scale-bg-crop to your notebook in advance.","metadata":{}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2022-08-16T12:49:27.220042Z","iopub.execute_input":"2022-08-16T12:49:27.222717Z","iopub.status.idle":"2022-08-16T12:49:28.454563Z","shell.execute_reply.started":"2022-08-16T12:49:27.222668Z","shell.execute_reply":"2022-08-16T12:49:28.453107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rm -r /kaggle/working/origin/train.csv","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport glob\nimport pandas as pd\nimport csv\n\nimg_dir = '/kaggle/input/stroke-blood-clot-origin-1k-scale-bg-crop/'\nannotations_file = '/kaggle/input/mayo-clinic-strip-ai/train.csv'\n\norigin_dir = '/kaggle/working/origin/'\n\nif not os.path.exists(origin_dir):\n    os.makedirs(origin_dir)\n    \norigin_annotation_file = os.path.join(origin_dir,\"train.csv\")\n\nls_imgs_png = glob.glob(os.path.join(img_dir, \"train_images\", \"*.png\"))\n    \nwith open(annotations_file, mode='r') as fr, open(origin_annotation_file, 'w') as fw:\n    reader = csv.reader(fr)\n    writer = csv.writer(fw)\n    \n    header = next(reader)\n    writer.writerow(header)\n    \n    for csv_line in reader:\n        for img_path in ls_imgs_png:\n            img_file_name = os.path.splitext(os.path.basename(img_path))[0]\n            if(csv_line[0] == img_file_name):\n                writer.writerow(csv_line)\n    \ndf_train = pd.read_csv(origin_annotation_file)\ndisplay(df_train)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T12:49:42.53535Z","iopub.execute_input":"2022-08-16T12:49:42.535835Z","iopub.status.idle":"2022-08-16T12:49:44.334849Z","shell.execute_reply.started":"2022-08-16T12:49:42.535792Z","shell.execute_reply":"2022-08-16T12:49:44.333881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-08-16T12:50:14.756381Z","iopub.execute_input":"2022-08-16T12:50:14.757171Z","iopub.status.idle":"2022-08-16T12:50:16.802273Z","shell.execute_reply.started":"2022-08-16T12:50:14.757132Z","shell.execute_reply":"2022-08-16T12:50:16.801117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset & DataLoader","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom torchvision.io import read_image\nfrom torchvision import datasets, models, transforms\nfrom torch.utils.data import Dataset, DataLoader\n\nclass ClotImageDataset(Dataset):\n    classes = ['CE','LAA']\n    def __init__(self, annotations_file, img_dir, transform=None, train = True, classes = classes):\n        self.img_labels = pd.read_csv(annotations_file)\n\n        if train:\n            self.img_dir = os.path.join(img_dir, \"train_images\")\n        else:\n            self.img_dir = os.path.join(img_dir, \"train_images\")\n            #self.img_dir = os.path.join(img_dir, \"other\")\n            \n        self.transform = transform\n        self.classes = classes\n\n    def __len__(self):\n        return len(self.img_labels)\n\n    def __getitem__(self, idx):\n        img_path = os.path.join(self.img_dir, self.img_labels.iloc[idx, 0] + \".png\")\n        image = read_image(img_path)\n\n        if self.transform:\n            image = self.transform(image)\n\n        # change graund truth labels to number label\n        label = self.classes.index(self.img_labels.iloc[idx, 4])\n\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2022-08-16T12:50:32.496107Z","iopub.execute_input":"2022-08-16T12:50:32.496704Z","iopub.status.idle":"2022-08-16T12:50:32.724519Z","shell.execute_reply.started":"2022-08-16T12:50:32.496668Z","shell.execute_reply":"2022-08-16T12:50:32.723534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = transforms.Compose([\n        transforms.ToPILImage(),\n        transforms.RandomResizedCrop(224),\n        transforms.RandomHorizontalFlip(),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n    ])\n\ntrain_dataset = ClotImageDataset(annotations_file = origin_annotation_file, img_dir = img_dir, transform = train_transforms, train = True)\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True, num_workers=2, drop_last=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T12:50:39.226324Z","iopub.execute_input":"2022-08-16T12:50:39.226712Z","iopub.status.idle":"2022-08-16T12:50:39.238241Z","shell.execute_reply.started":"2022-08-16T12:50:39.226681Z","shell.execute_reply":"2022-08-16T12:50:39.237213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define model","metadata":{}},{"cell_type":"code","source":"from torchvision.models import resnet50\nimport torch.nn as nn\n\nmodel = resnet34(pretrained = True)\nmodel.fc = nn.Linear(2048, len(train_dataset.classes))\nnn.init.normal_(model.fc.weight, mean=0, std=5e-3)\nmodel.fc.bias.data.fill_(0.01)\nmodel.to(device)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision.models import resnet34\nimport torch.nn as nn\n\nmodel = resnet34(pretrained = True)\nmodel.fc = nn.Linear(512, len(train_dataset.classes))\nnn.init.normal_(model.fc.weight, mean=0, std=5e-3)\nmodel.fc.bias.data.fill_(0.01)\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T13:35:54.940193Z","iopub.execute_input":"2022-08-16T13:35:54.940622Z","iopub.status.idle":"2022-08-16T13:35:55.456508Z","shell.execute_reply.started":"2022-08-16T13:35:54.940587Z","shell.execute_reply":"2022-08-16T13:35:55.45551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define optimization ","metadata":{}},{"cell_type":"code","source":"from torch import optim\nlr = 1e-2\nmomentum = 0.9\n\nparams = []\nfor name,param in model.named_parameters():\n    if param.requires_grad == True:\n        params.append(param)\n        \noptimizer = optim.SGD([\n        {'params':  params[:-3], 'lr':1.0*lr},\n        {'params':  params[-2:], 'lr':10.0*lr}\n    ], lr=lr, momentum=momentum)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T13:35:59.139236Z","iopub.execute_input":"2022-08-16T13:35:59.139638Z","iopub.status.idle":"2022-08-16T13:35:59.147173Z","shell.execute_reply.started":"2022-08-16T13:35:59.139603Z","shell.execute_reply":"2022-08-16T13:35:59.14615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Difine Loss Function","metadata":{}},{"cell_type":"code","source":"ce_loss = torch.nn.CrossEntropyLoss()\nce_loss = ce_loss.to(device)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T13:36:01.702177Z","iopub.execute_input":"2022-08-16T13:36:01.703141Z","iopub.status.idle":"2022-08-16T13:36:01.708287Z","shell.execute_reply.started":"2022-08-16T13:36:01.703103Z","shell.execute_reply":"2022-08-16T13:36:01.70709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Training & Testing codes","metadata":{}},{"cell_type":"code","source":"def train(model, train_loader, optimizer, ce_loss):\n    losses = AverageMeter()\n    top1 = AverageMeter()\n    model.train()\n    \n    for iteration in range(len(train_loader)):\n        img, label = next(iter(train_loader))\n        img ,label = img.to(device), label.to(device)\n        optimizer.zero_grad()\n\n        out = model(img)\n        loss = ce_loss(out, label)\n\n        prec1, _ = accuracy(out, label,topk=(1,2))\n        losses.update(loss.item(), img.size(0))\n        top1.update(prec1.item(), img.size(0))\n        loss.backward()\n        optimizer.step()\n        \n    print('Cls:{cls_losses.val:.4f}({cls_losses.avg:.4f})  '\n                'prec@1:{top1.val:.2f}({top1.avg:.2f})  '\n                .format(cls_losses=losses,top1=top1))\n        \nclass AverageMeter(object):\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val   = 0\n        self.avg   = 0\n        self.sum   = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val   = val\n        self.sum   += val * n\n        self.count += n\n        self.avg   = self.sum / self.count\n\ndef accuracy(output, target, topk=(1,)):\n    with torch.no_grad():\n        maxk = max(topk)\n        batch_size = target.size(0)\n\n        _, pred = output.topk(maxk, 1, True, True)\n        pred = pred.t()\n        correct = pred.eq(target[None])\n\n        res = []\n        for k in topk:\n            correct_k = correct[:k].flatten().sum(dtype=torch.float32)\n            res.append(correct_k * (100.0 / batch_size))\n        return res\n        ","metadata":{"execution":{"iopub.status.busy":"2022-08-16T13:36:35.042426Z","iopub.execute_input":"2022-08-16T13:36:35.04319Z","iopub.status.idle":"2022-08-16T13:36:35.055287Z","shell.execute_reply.started":"2022-08-16T13:36:35.043152Z","shell.execute_reply":"2022-08-16T13:36:35.054229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define main code","metadata":{}},{"cell_type":"code","source":"def main(num_epochs):\n    for epoch in range(1, num_epochs+1):\n        print('-----------------------')\n        print('Epoch {}/{}'.format(epoch,num_epochs))\n\n        train(model, train_loader, optimizer, ce_loss)\n        ","metadata":{"execution":{"iopub.status.busy":"2022-08-16T13:36:06.966475Z","iopub.execute_input":"2022-08-16T13:36:06.967647Z","iopub.status.idle":"2022-08-16T13:36:06.973389Z","shell.execute_reply.started":"2022-08-16T13:36:06.967602Z","shell.execute_reply":"2022-08-16T13:36:06.972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Experiment all code","metadata":{}},{"cell_type":"code","source":"main(10)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T13:36:39.336577Z","iopub.execute_input":"2022-08-16T13:36:39.337175Z","iopub.status.idle":"2022-08-16T13:47:45.615896Z","shell.execute_reply.started":"2022-08-16T13:36:39.337141Z","shell.execute_reply":"2022-08-16T13:47:45.611911Z"},"trusted":true},"execution_count":null,"outputs":[]}]}