{"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 libauc","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:02:06.328593Z","iopub.execute_input":"2022-03-02T20:02:06.331446Z","iopub.status.idle":"2022-03-02T20:02:34.699815Z","shell.execute_reply.started":"2022-03-02T20:02:06.331326Z","shell.execute_reply":"2022-03-02T20:02:34.698952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from libauc.losses import AUCMLoss, CrossEntropyLoss\nfrom libauc.optimizers import PESG, Adam\n\nimport torch \nimport numpy as np\nimport torch.utils.data as dt\nfrom sklearn.metrics import roc_auc_score","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:02:34.702979Z","iopub.execute_input":"2022-03-02T20:02:34.703231Z","iopub.status.idle":"2022-03-02T20:02:37.129779Z","shell.execute_reply.started":"2022-03-02T20:02:34.703185Z","shell.execute_reply":"2022-03-02T20:02:37.128993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nfrom torch.utils.data import Dataset\nfrom torchvision import transforms\nfrom PIL import Image\n\nclass ProteinDataset(Dataset):\n    def __init__(self, images_path:str, labels_path:str, image_size=512, n_samples_to_load=None):\n        self.labels = pd.read_csv(labels_path, nrows=n_samples_to_load)\n        self.img_path = images_path \n        self.preprocess = transforms.Compose([            \n            transforms.Resize((image_size, image_size)),\n            transforms.RandomVerticalFlip(),\n            transforms.RandomHorizontalFlip(),      \n            transforms.RandomRotation(90),      \n            transforms.ToTensor(),            \n            # transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        ])\n\n\n    def __len__(self):\n        return len(self.labels)\n\n\n    def __getitem__(self, idx):        \n        id = self.labels.iloc[idx]['Id']\n        label = transform_label_from_string(self.labels.iloc[idx]['Target']) \n        red    = Image.open(f'{self.img_path}{id}_red.png')\n        blue   = Image.open(f'{self.img_path}{id}_blue.png')\n        green  = Image.open(f'{self.img_path}{id}_green.png')        \n        yellow = Image.open(f'{self.img_path}{id}_yellow.png')\n        img = Image.merge('RGBX',[red,blue,green,yellow])\n        img = self.preprocess(img)        \n        return img, label\n\ndef transform_label_from_string(label:str):\n    lab = label.split(' ')\n    new_lab = np.zeros(28)\n    for l in lab:\n        position = int(l)\n        new_lab[position] = 1\n    return new_lab\n\ndef transform_label_to_string(label):\n    new_lab = []\n    for i,lab in enumerate(label):\n        if lab>0.5:\n            new_lab.insert(0,str(i))\n    str_label = ' '.join(new_lab)\n    return str_label\n    \n\n    \n        \n","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:02:37.131255Z","iopub.execute_input":"2022-03-02T20:02:37.131511Z","iopub.status.idle":"2022-03-02T20:02:37.490091Z","shell.execute_reply.started":"2022-03-02T20:02:37.131475Z","shell.execute_reply":"2022-03-02T20:02:37.489284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F \n\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n        \n    def forward(self, input, target):\n        if not (target.size() == input.size()):\n            raise ValueError(\"Target size ({}) must be the same as input size ({})\"\n                             .format(target.size(), input.size()))\n\n        max_val = (-input).clamp(min=0)\n        loss = input - input * target + max_val + \\\n            ((-max_val).exp() + (-input - max_val).exp()).log()\n\n        invprobs = F.logsigmoid(-input * (target * 2.0 - 1.0))\n        loss = (invprobs * self.gamma).exp() * loss\n        \n        return loss.sum(dim=1).mean()\n\nclass BinaryDiceLoss(nn.Module):\n    \"\"\"Dice loss of binary class\n    Args:\n        smooth: A float number to smooth loss, and avoid NaN error, default: 1\n        p: Denominator value: \\sum{x^p} + \\sum{y^p}, default: 2\n        predict: A tensor of shape [N, *]\n        target: A tensor of shape same with predict\n        reduction: Reduction method to apply, return mean over batch if 'mean',\n            return sum if 'sum', return a tensor of shape [N,] if 'none'\n    Returns:\n        Loss tensor according to arg reduction\n    Raise:\n        Exception if unexpected reduction\n    \"\"\"\n    def __init__(self, smooth=1, p=2, reduction='mean'):\n        super(BinaryDiceLoss, self).__init__()\n        self.smooth = smooth\n        self.p = p\n        self.reduction = reduction\n\n    def forward(self, predict, target):\n        assert predict.shape[0] == target.shape[0], \"predict & target batch size don't match\"\n        predict = predict.contiguous().view(predict.shape[0], -1)\n        target = target.contiguous().view(target.shape[0], -1)\n\n        num = torch.sum(torch.mul(predict, target), dim=1) + self.smooth\n        den = torch.sum(predict.pow(self.p) + target.pow(self.p), dim=1) + self.smooth\n\n        loss = 1 - num / den\n\n        if self.reduction == 'mean':\n            return loss.mean()\n        elif self.reduction == 'sum':\n            return loss.sum()\n        elif self.reduction == 'none':\n            return loss\n        else:\n            raise Exception('Unexpected reduction {}'.format(self.reduction))\n\nclass DiceLoss(nn.Module):\n    \"\"\"Dice loss, need one hot encode input\n    Args:\n        weight: An array of shape [num_classes,]\n        ignore_index: class index to ignore\n        predict: A tensor of shape [N, C, *]\n        target: A tensor of same shape with predict\n        other args pass to BinaryDiceLoss\n    Return:\n        same as BinaryDiceLoss\n    \"\"\"\n    def __init__(self, weight=None, ignore_index=None, **kwargs):\n        super(DiceLoss, self).__init__()\n        self.kwargs = kwargs\n        self.weight = weight\n        self.ignore_index = ignore_index\n\n    def forward(self, predict, target):\n        assert predict.shape == target.shape, 'predict & target shape do not match'\n        dice = BinaryDiceLoss(**self.kwargs)\n        total_loss = 0\n        predict = F.softmax(predict, dim=1)\n\n        for i in range(target.shape[1]):\n            if i != self.ignore_index:\n                dice_loss = dice(predict[:, i], target[:, i])\n                if self.weight is not None:\n                    assert self.weight.shape[0] == target.shape[1], \\\n                        'Expect weight shape [{}], get[{}]'.format(target.shape[1], self.weight.shape[0])\n                    dice_loss *= self.weights[i]\n                total_loss += dice_loss\n\n        return total_loss/target.shape[1]\n\n\ndef diceCoeffv2(pred, gt, eps=1e-5, activation='sigmoid'):\n    r\"\"\" computational formula：\n        dice = (2 * tp) / (2 * tp + fp + fn)\n    \"\"\"\n \n    if activation is None or activation == \"none\":\n        activation_fn = lambda x: x\n    elif activation == \"sigmoid\":\n        activation_fn = nn.Sigmoid()\n    elif activation == \"softmax2d\":\n        activation_fn = nn.Softmax2d()\n    else:\n        raise NotImplementedError # («Активация реализована для работы функции активации сигмоида и softmax2d»)\n \n    pred = activation_fn(pred)\n \n    N = gt.size(0)\n    pred_flat = pred.view(N, -1)\n    gt_flat = gt.view(N, -1)\n \n    tp = torch.sum(gt_flat * pred_flat, dim=1)\n    fp = torch.sum(pred_flat, dim=1) - tp\n    fn = torch.sum(gt_flat, dim=1) - tp\n    loss = (2 * tp + eps) / (2 * tp + fp + fn + eps)\n    return loss.sum() / N\n\n\nclass DiceLossV2(nn.Module):\n    __name__ = 'dice_loss'\n \n    def __init__(self, activation='sigmoid'):\n        super(DiceLossV2, self).__init__()\n        self.activation = activation\n \n    def forward(self, y_pr, y_gt):\n        return 1 - diceCoeffv2(y_pr, y_gt, activation=self.activation)\n\ndef dice_loss(input, target):\n    input = torch.sigmoid(input)\n    smooth = 1.0\n\n    iflat = input.view(-1)\n    tflat = target.view(-1)\n    intersection = (iflat * tflat).sum()\n    \n    return ((2.0 * intersection + smooth) / (iflat.sum() + tflat.sum() + smooth))\n\n\nclass MixedLoss(nn.Module):\n    def __init__(self, alpha, gamma):\n        super().__init__()\n        self.alpha = alpha\n        self.focal = FocalLoss(gamma)\n        \n    def forward(self, input, target):\n        loss = self.alpha*self.focal(input, target) - torch.log(dice_loss(input, target))\n        return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:02:37.492530Z","iopub.execute_input":"2022-03-02T20:02:37.493041Z","iopub.status.idle":"2022-03-02T20:02:37.538341Z","shell.execute_reply.started":"2022-03-02T20:02:37.493003Z","shell.execute_reply":"2022-03-02T20:02:37.537179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_all_seeds(SEED):\n    # REPRODUCIBILITY\n    torch.manual_seed(SEED)\n    np.random.seed(SEED)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:02:37.540189Z","iopub.execute_input":"2022-03-02T20:02:37.540814Z","iopub.status.idle":"2022-03-02T20:02:37.546074Z","shell.execute_reply.started":"2022-03-02T20:02:37.540776Z","shell.execute_reply":"2022-03-02T20:02:37.545260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_DIR = '../input/human-protein-atlas-image-classification/train/'\nTEST_DIR = '../input/human-protein-atlas-image-classification/test/'\nLABELS = '../input/human-protein-atlas-image-classification/train.csv'\ndataset = ProteinDataset(TRAIN_DIR, LABELS, image_size=512)","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:02:37.547752Z","iopub.execute_input":"2022-03-02T20:02:37.548337Z","iopub.status.idle":"2022-03-02T20:02:37.622024Z","shell.execute_reply.started":"2022-03-02T20:02:37.548302Z","shell.execute_reply":"2022-03-02T20:02:37.621265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install timm","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:02:37.626216Z","iopub.execute_input":"2022-03-02T20:02:37.628017Z","iopub.status.idle":"2022-03-02T20:02:46.115616Z","shell.execute_reply.started":"2022-03-02T20:02:37.627981Z","shell.execute_reply":"2022-03-02T20:02:46.114730Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nfrom tqdm import tqdm\n# dataloader\n# Create dataloaders\nbatch_size = 12\ndataset_len = dataset.__len__()\ntrain_size = int(dataset_len*0.8)\nif train_size % batch_size==1:\n        train_size += 1\n\nval_size = dataset_len - train_size\ntrain_set, val_set = torch.utils.data.random_split(dataset, [train_size, val_size])\n    \ntrainloader = dt.DataLoader(train_set, batch_size=batch_size)\ntestloader = dt.DataLoader(val_set, batch_size=batch_size)   \n\n# paramaters\nSEED = 123\nlr = 1e-4\nweight_decay = 1e-5\n\n# model\nset_all_seeds(SEED)\nmodel= timm.create_model('efficientnet_b4',pretrained=True,num_classes=28,in_chans=4)\n# model.load_state_dict(torch.load('./modified_pretrained_model.pth'))\nmodel = model.cuda()\n\n# define loss & optimizer\ncriterion = CrossEntropyLoss()\noptimizer = Adam(model.parameters(), lr=lr, weight_decay=weight_decay)\n\n# training\nbest_val_auc = 0 \nfor epoch in range(1):\n    model.train()\n    batch_index = 0\n    epoch_loss = []\n    for data in tqdm(trainloader):\n        batch_index +=1\n        train_data, train_labels = data\n        train_data, train_labels  = train_data.cuda(), train_labels.cuda()\n        y_pred = model(train_data)\n        loss = criterion(y_pred, train_labels)\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n          \n        epoch_loss.append(loss.item())\n\n        # validation  \n    \n    model.eval()\n    with torch.no_grad():    \n        test_pred = []\n        test_true = [] \n        for data in tqdm(testloader):\n                    test_data, test_labels = data\n                    test_data = test_data.cuda()\n                    y_pred = model(test_data)\n                    test_pred.append(y_pred.cpu().detach().numpy())\n                    test_true.append(test_labels.numpy())\n            \n        test_true = np.concatenate(test_true)\n        test_pred = np.concatenate(test_pred)\n        val_auc_mean =  roc_auc_score(test_true, test_pred) \n        model.train()\n\n        if best_val_auc < val_auc_mean:\n            best_val_auc = val_auc_mean\n            torch.save(model.state_dict(), 'modified_pretrained_model_1.pth')\n\n        print ('Epoch=%s, BatchID=%s, Val_AUC=%.4f, Best_Val_AUC=%.4f'%(epoch, batch_index, val_auc_mean, best_val_auc ))\n    print (f\"Epoch loss {np.mean(epoch_loss)}\")","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:03:42.669829Z","iopub.execute_input":"2022-03-02T20:03:42.670097Z","iopub.status.idle":"2022-03-02T20:52:34.872130Z","shell.execute_reply.started":"2022-03-02T20:03:42.670067Z","shell.execute_reply":"2022-03-02T20:52:34.871353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('../input/human-protein-atlas-image-classification/train.csv')\ntrain_df['split_targ'] = train_df['Target'].str.split(' ')\n\nfor i in range(28):\n    train_df[f'{i}'] =train_df['split_targ'].apply(lambda x:  str(i) in x)\ncounts = train_df.sum().drop(['Id','Target','split_targ'])\nless_than_100 = counts[counts<100].index\nless_than_250 = counts[(counts<250) & (counts>=100)].index\nless_than_500 = counts[(counts<500) & (counts>=250)].index\nless_than_100, less_than_250, less_than_500\n","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:52:47.183707Z","iopub.execute_input":"2022-03-02T20:52:47.183969Z","iopub.status.idle":"2022-03-02T20:52:50.824240Z","shell.execute_reply.started":"2022-03-02T20:52:47.183941Z","shell.execute_reply":"2022-03-02T20:52:50.823510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_df.shape)\nupsampled = train_df\nfor i in range(5):\n    upsampled = pd.concat([upsampled,train_df[(train_df['8']==True) | (train_df['9']==True) | (train_df['10']==True) | (train_df['15']==True) | (train_df['27']==True) ]],ignore_index=True)\nfor i in range(3):\n    upsampled = pd.concat([upsampled,train_df[(train_df['17']==True) | (train_df['20']==True)]],ignore_index=True)    \nfor i in range(2):\n    upsampled = pd.concat([upsampled,train_df[(train_df['24']==True) | (train_df['26']==True)]],ignore_index=True) \n\nprint(upsampled.shape)","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:53:02.159927Z","iopub.execute_input":"2022-03-02T20:53:02.160184Z","iopub.status.idle":"2022-03-02T20:53:02.201993Z","shell.execute_reply.started":"2022-03-02T20:53:02.160154Z","shell.execute_reply":"2022-03-02T20:53:02.201141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"upsampled = upsampled[['Id','Target']]\n\nupsampled =upsampled.sample(frac=1).reset_index()\nupsampled.to_csv('train_upsampled.csv')\nupsampled","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:53:07.742072Z","iopub.execute_input":"2022-03-02T20:53:07.742632Z","iopub.status.idle":"2022-03-02T20:53:07.893862Z","shell.execute_reply.started":"2022-03-02T20:53:07.742595Z","shell.execute_reply":"2022-03-02T20:53:07.893039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LABELS = './train_upsampled.csv'\ndataset = ProteinDataset(TRAIN_DIR, LABELS, image_size=512)\ndataset_len = dataset.__len__()\ntrain_size = int(dataset_len*0.8)\nif train_size % batch_size==1:\n        train_size += 1\n\nval_size = dataset_len - train_size\ntrain_set, val_set = torch.utils.data.random_split(dataset, [train_size, val_size])\n    \ntrainloader = dt.DataLoader(train_set, batch_size=batch_size)\ntestloader = dt.DataLoader(val_set, batch_size=batch_size) ","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:53:11.064962Z","iopub.execute_input":"2022-03-02T20:53:11.065667Z","iopub.status.idle":"2022-03-02T20:53:11.108842Z","shell.execute_reply.started":"2022-03-02T20:53:11.065631Z","shell.execute_reply":"2022-03-02T20:53:11.108063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = model.cuda()\n\n# define loss & optimizer\ncriterion = CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-5, weight_decay=weight_decay)\n\n# training\nbest_val_auc = 0 \nfor epoch in range(2):\n    model.train()\n    batch_index = 0\n    epoch_loss = []\n    for data in tqdm(trainloader):\n        batch_index +=1\n        train_data, train_labels = data\n        train_data, train_labels  = train_data.cuda(), train_labels.cuda()\n        y_pred = model(train_data)\n        loss = criterion(y_pred, train_labels)\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n          \n        epoch_loss.append(loss.item())\n\n        # validation  \n    \n    model.eval()\n    with torch.no_grad():    \n        test_pred = []\n        test_true = [] \n        for data in tqdm(testloader):\n                    test_data, test_labels = data\n                    test_data = test_data.cuda()\n                    y_pred = model(test_data)\n                    test_pred.append(y_pred.cpu().detach().numpy())\n                    test_true.append(test_labels.numpy())\n            \n        test_true = np.concatenate(test_true)\n        test_pred = np.concatenate(test_pred)\n        val_auc_mean =  roc_auc_score(test_true, test_pred) \n        model.train()\n\n        if best_val_auc < val_auc_mean:\n            best_val_auc = val_auc_mean\n            torch.save(model.state_dict(), f'modified_pretrained_model_2_{epoch}.pth')\n\n        print ('Epoch=%s, BatchID=%s, Val_AUC=%.4f, Best_Val_AUC=%.4f'%(epoch, batch_index, val_auc_mean, best_val_auc ))\n    print (f\"Epoch loss {np.mean(epoch_loss)}\")","metadata":{"execution":{"iopub.status.busy":"2022-03-02T20:53:15.768630Z","iopub.execute_input":"2022-03-02T20:53:15.768886Z","iopub.status.idle":"2022-03-02T22:26:40.138090Z","shell.execute_reply.started":"2022-03-02T20:53:15.768857Z","shell.execute_reply":"2022-03-02T22:26:40.137365Z"},"trusted":true},"execution_count":null,"outputs":[]}]}