{"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":"# # This Python 3 environment comes with many helpful analytics libraries installed\n# # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# # For example, here's several helpful packages to load\n\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-06-07T14:39:44.589639Z","iopub.execute_input":"2022-06-07T14:39:44.589996Z","iopub.status.idle":"2022-06-07T14:39:44.612822Z","shell.execute_reply.started":"2022-06-07T14:39:44.589924Z","shell.execute_reply":"2022-06-07T14:39:44.612184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport random\nimport os\nimport cv2\nimport gc\nimport warnings\nfrom sklearn.metrics import f1_score\nfrom sklearn.exceptions import UndefinedMetricWarning\nimport scipy.optimize as opt\nfrom collections import defaultdict, Counter\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport torch\nfrom torch import nn, optim\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models import resnet34\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:44.614235Z","iopub.execute_input":"2022-06-07T14:39:44.614907Z","iopub.status.idle":"2022-06-07T14:39:48.496543Z","shell.execute_reply.started":"2022-06-07T14:39:44.614872Z","shell.execute_reply":"2022-06-07T14:39:48.495487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#DIR_PATHS\nTRAIN = '../input/hpa-single-cell-image-classification/train'\nTEST = '../input/hpa-single-cell-image-classification/test'\nLABELS = '../input/hpa-single-cell-image-classification/train.csv'\nSUBMIT = '../input/hpa-single-cell-image-classification/sample_submission.csv'","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:48.501956Z","iopub.execute_input":"2022-06-07T14:39:48.503964Z","iopub.status.idle":"2022-06-07T14:39:48.514763Z","shell.execute_reply.started":"2022-06-07T14:39:48.503924Z","shell.execute_reply":"2022-06-07T14:39:48.512477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Labels for dataset\nLABELS_DICT = {\n    0: 'Nucleoplasm',\n    1: 'Nuclear membrane',\n    2: 'Nucleoli',\n    3: 'Nucleoli fibrilar center',\n    4: 'Nucler speckles',\n    5: 'Nuclear bodies',\n    6: 'Endoplasmic reticulum',\n    7: 'Golgi apparatus',\n    8: 'Intermediate filaments',\n    9: 'Actin filaments',\n    10: 'Microtubules',\n    11: 'Mitotic spindle',\n    12: 'Centrosome',\n    13: 'Plasma membrane',\n    14: 'Mitochondria',\n    15: 'Aggresome',\n    16: 'Cytosol',\n    17: 'Vesicles and punctate cytosolic patterns',\n    18: 'Negative'\n}","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:48.519776Z","iopub.execute_input":"2022-06-07T14:39:48.520291Z","iopub.status.idle":"2022-06-07T14:39:48.528649Z","shell.execute_reply.started":"2022-06-07T14:39:48.520207Z","shell.execute_reply":"2022-06-07T14:39:48.527643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(LABELS)\ndf.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:48.530195Z","iopub.execute_input":"2022-06-07T14:39:48.530832Z","iopub.status.idle":"2022-06-07T14:39:48.699988Z","shell.execute_reply.started":"2022-06-07T14:39:48.530791Z","shell.execute_reply":"2022-06-07T14:39:48.699047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.sample()","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:48.70165Z","iopub.execute_input":"2022-06-07T14:39:48.701922Z","iopub.status.idle":"2022-06-07T14:39:48.739205Z","shell.execute_reply.started":"2022-06-07T14:39:48.701889Z","shell.execute_reply":"2022-06-07T14:39:48.73848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv(SUBMIT)\nsub_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:48.7406Z","iopub.execute_input":"2022-06-07T14:39:48.741138Z","iopub.status.idle":"2022-06-07T14:39:48.793517Z","shell.execute_reply.started":"2022-06-07T14:39:48.741101Z","shell.execute_reply":"2022-06-07T14:39:48.792691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Count for each label in dataset\n\nN_CLASSES = len(LABELS_DICT)\ncls_counts = Counter(cls for classes in df['Label'].str.split('|') for cls in classes)\ncounts_x = [i[1] for i in cls_counts.most_common(N_CLASSES)]\ncounts_y = [LABELS_DICT[int(i[0])] for i in cls_counts.most_common(N_CLASSES)]\nplt.figure(figsize=(8,8))\nsns.barplot(y=counts_y, x=counts_x)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:48.794885Z","iopub.execute_input":"2022-06-07T14:39:48.795119Z","iopub.status.idle":"2022-06-07T14:39:49.47895Z","shell.execute_reply.started":"2022-06-07T14:39:48.795088Z","shell.execute_reply":"2022-06-07T14:39:49.477687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_folds = 10\nfold_cls_counts = defaultdict(int)\n# print(fold_cls_counts)\nfolds = [-1] * len(df)\nfor item in tqdm(df.sample(frac=1, random_state=42).itertuples(),total=len(df)):\n    cls = min(item.Label.split('|'), key=lambda cls: cls_counts[cls])\n    fold_counts = [(f, fold_cls_counts[f, cls]) for f in range(n_folds)]\n    min_count = min([count for _, count in fold_counts])\n    random.seed(item.Index)\n#     print(item.Index)\n    fold = random.choice([f for f, count in fold_counts if count == min_count])\n    folds[item.Index] = fold\n    for cls in item.Label.split():\n        fold_cls_counts[fold, cls] += 1\nprint(fold_cls_counts)\ndf['fold'] = folds","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:49.481765Z","iopub.execute_input":"2022-06-07T14:39:49.482418Z","iopub.status.idle":"2022-06-07T14:39:50.075266Z","shell.execute_reply.started":"2022-06-07T14:39:49.482381Z","shell.execute_reply":"2022-06-07T14:39:50.074546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:50.076391Z","iopub.execute_input":"2022-06-07T14:39:50.07702Z","iopub.status.idle":"2022-06-07T14:39:50.09186Z","shell.execute_reply.started":"2022-06-07T14:39:50.076987Z","shell.execute_reply":"2022-06-07T14:39:50.090926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df.head(50)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:50.094376Z","iopub.execute_input":"2022-06-07T14:39:50.094565Z","iopub.status.idle":"2022-06-07T14:39:50.102435Z","shell.execute_reply.started":"2022-06-07T14:39:50.094542Z","shell.execute_reply":"2022-06-07T14:39:50.101695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_idx = 0\ntrain_df = df[df['fold']!=valid_idx][['ID', 'Label']].reset_index(drop=True)\nvalid_df = df[df['fold']==valid_idx][['ID', 'Label']].reset_index(drop=True)\nprint('There are {} samples in the training set.'.format(len(train_df)))\nprint('There are {} samples in the validation set.'.format(len(valid_df)))","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:50.104064Z","iopub.execute_input":"2022-06-07T14:39:50.104281Z","iopub.status.idle":"2022-06-07T14:39:50.119478Z","shell.execute_reply.started":"2022-06-07T14:39:50.104255Z","shell.execute_reply":"2022-06-07T14:39:50.118716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:50.120732Z","iopub.execute_input":"2022-06-07T14:39:50.12099Z","iopub.status.idle":"2022-06-07T14:39:50.129308Z","shell.execute_reply.started":"2022-06-07T14:39:50.120957Z","shell.execute_reply":"2022-06-07T14:39:50.128474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_df.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:50.130758Z","iopub.execute_input":"2022-06-07T14:39:50.13123Z","iopub.status.idle":"2022-06-07T14:39:50.143308Z","shell.execute_reply.started":"2022-06-07T14:39:50.131147Z","shell.execute_reply":"2022-06-07T14:39:50.142435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def open_rgby(path, id, size=None):\n    colors = ['red', 'green', 'blue', 'yellow']\n    flags = cv2.IMREAD_GRAYSCALE\n    if size is None:\n        img = [cv2.imread(os.path.join(path, id + '_' +color+'.png'), flags).astype(np.float32)/255\n              for color in colors]\n    else:\n        img = []\n        for color in colors:\n            src_img = cv2.imread(os.path.join(path, id + '_' +color+'.png'), flags)\n            tar_img = cv2.resize(src_img, (size, size), interpolation = cv2.INTER_CUBIC).astype(np.float32)/255\n            img.append(tar_img)\n    return np.stack(img, axis = 0)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:50.144676Z","iopub.execute_input":"2022-06-07T14:39:50.145009Z","iopub.status.idle":"2022-06-07T14:39:50.15615Z","shell.execute_reply.started":"2022-06-07T14:39:50.144973Z","shell.execute_reply":"2022-06-07T14:39:50.155461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_image_row(image, subax, title):\n    subax[0].imshow(np.stack(image, axis=-1)[:,:,:3])\n    subax[0].set_title('ID: '+title)\n    subax[1].imshow(image[0], cmap=\"Reds\")\n    subax[1].set_title(\"Red Channel\")\n    subax[2].imshow(image[1], cmap=\"Greens\")\n    subax[2].set_title(\"Green Channel\")\n    subax[3].imshow(image[2], cmap=\"Blues\")\n    subax[3].set_title(\"Blue Channel\")\n    subax[4].imshow(image[3], cmap=\"Oranges\")\n    subax[4].set_title(\"Yellow Channel\")\n    return subax","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:50.157616Z","iopub.execute_input":"2022-06-07T14:39:50.158298Z","iopub.status.idle":"2022-06-07T14:39:50.168121Z","shell.execute_reply.started":"2022-06-07T14:39:50.158261Z","shell.execute_reply":"2022-06-07T14:39:50.167403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_SAMPLES = 5\nsamples = np.array(train_df.sample(N_SAMPLES)['ID'])\n\nfig, ax = plt.subplots(N_SAMPLES,5,figsize=(20,5*N_SAMPLES))\n\nif ax.shape == (N_SAMPLES,):\n    ax = ax.reshape(1,-1)\n#     print(ax)\n# else :\n#     print(ax)\nfor n in range(N_SAMPLES):\n    make_image_row(open_rgby(TRAIN, samples[n]), ax[n], samples[n].split('-')[0])","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:39:50.169545Z","iopub.execute_input":"2022-06-07T14:39:50.169822Z","iopub.status.idle":"2022-06-07T14:40:14.735101Z","shell.execute_reply.started":"2022-06-07T14:39:50.169788Z","shell.execute_reply":"2022-06-07T14:40:14.734303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"After the read sample images function has been defined, build dataset interface for Pytorch Model.\nAs this is RGBY images (4 channels)","metadata":{}},{"cell_type":"code","source":"# samples = np.array(train_df.sample(N_SAMPLES)['ID'])","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.736162Z","iopub.execute_input":"2022-06-07T14:40:14.736424Z","iopub.status.idle":"2022-06-07T14:40:14.740267Z","shell.execute_reply.started":"2022-06-07T14:40:14.736385Z","shell.execute_reply":"2022-06-07T14:40:14.739625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# samples[2].split('-')[0]\n# open_rgby(TRAIN,samples[2]).itemsize","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.741255Z","iopub.execute_input":"2022-06-07T14:40:14.741716Z","iopub.status.idle":"2022-06-07T14:40:14.750618Z","shell.execute_reply.started":"2022-06-07T14:40:14.741683Z","shell.execute_reply":"2022-06-07T14:40:14.749703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def open_rgby2(path, id, size=256):\n    R = cv2.resize(cv2.imread(path + '/' + id + '_red.png', cv2.IMREAD_GRAYSCALE),(size,size),\n                   interpolation=cv2.INTER_CUBIC).astype(np.float32)/255\n    G = cv2.resize(cv2.imread(path + '/' + id + '_green.png', cv2.IMREAD_GRAYSCALE),(size,size),\n                   interpolation=cv2.INTER_CUBIC).astype(np.float32)/255\n    B = cv2.resize(cv2.imread(path + '/' + id + '_blue.png', cv2.IMREAD_GRAYSCALE),(size,size),\n                   interpolation=cv2.INTER_CUBIC).astype(np.float32)/255\n    Y = cv2.resize(cv2.imread(path + '/' + id + '_yellow.png', cv2.IMREAD_GRAYSCALE),(size,size),\n                   interpolation=cv2.INTER_CUBIC).astype(np.float32)/255\n    img = np.stack((R*2/3+Y/3,\n                    G*2/3+ Y/3,\n                    B*2/3+ Y/3), axis=0)\n    return img","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.752063Z","iopub.execute_input":"2022-06-07T14:40:14.753019Z","iopub.status.idle":"2022-06-07T14:40:14.764239Z","shell.execute_reply.started":"2022-06-07T14:40:14.752976Z","shell.execute_reply":"2022-06-07T14:40:14.763241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# R = cv2.imread(TEST + '/0040581b-f1f2-4fbe-b043-b6bfea5404bb_red.png', cv2.IMREAD_GRAYSCALE)\n# G = cv2.imread(TEST + '/0040581b-f1f2-4fbe-b043-b6bfea5404bb_green.png', cv2.IMREAD_GRAYSCALE)\n# B = cv2.imread(TEST + '/0040581b-f1f2-4fbe-b043-b6bfea5404bb_blue.png', cv2.IMREAD_GRAYSCALE)\n# Y = cv2.imread(TEST + '/0040581b-f1f2-4fbe-b043-b6bfea5404bb_yellow.png', cv2.IMREAD_GRAYSCALE)\n","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.765365Z","iopub.execute_input":"2022-06-07T14:40:14.765746Z","iopub.status.idle":"2022-06-07T14:40:14.776376Z","shell.execute_reply.started":"2022-06-07T14:40:14.765711Z","shell.execute_reply":"2022-06-07T14:40:14.775702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# img = np.stack((R*2/3+Y/3,\n#                G*2/3+ Y/3,\n#                B*2/3+ Y/3), -1)\n# img = cv2.resize(img, (256,256))\n# img = np.divide(img, 255)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.777461Z","iopub.execute_input":"2022-06-07T14:40:14.77831Z","iopub.status.idle":"2022-06-07T14:40:14.785154Z","shell.execute_reply.started":"2022-06-07T14:40:14.778247Z","shell.execute_reply":"2022-06-07T14:40:14.784459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def make_image_row2(image, subax, title):\n#     subax[0].imshow(np.stack(image, axis=-1)[:,:,:3])\n#     subax[0].set_title('Id: '+title)\n#     subax[1].imshow(image[0], cmap=\"Reds\")\n#     subax[1].set_title(\"Red Channel\")\n#     subax[2].imshow(image[1], cmap=\"Greens\")\n#     subax[2].set_title(\"Green Channel\")\n#     subax[3].imshow(image[2], cmap=\"Blues\")\n#     subax[3].set_title(\"Blue Channel\")\n#     return subax","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.786902Z","iopub.execute_input":"2022-06-07T14:40:14.787455Z","iopub.status.idle":"2022-06-07T14:40:14.794703Z","shell.execute_reply.started":"2022-06-07T14:40:14.787418Z","shell.execute_reply":"2022-06-07T14:40:14.793632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sample = np.array(train_df.sample()['ID'])\n# sample[0]","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.800701Z","iopub.execute_input":"2022-06-07T14:40:14.801153Z","iopub.status.idle":"2022-06-07T14:40:14.804525Z","shell.execute_reply.started":"2022-06-07T14:40:14.801118Z","shell.execute_reply":"2022-06-07T14:40:14.803577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fig, ax = plt.subplots(1,4,figsize=(20,5))\n# if ax.shape == (1,):\n#     ax = ax.reshape(1,-1)\n# make_image_row(img, ax[n], samples[n].split('-')[0])","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.806353Z","iopub.execute_input":"2022-06-07T14:40:14.806945Z","iopub.status.idle":"2022-06-07T14:40:14.814132Z","shell.execute_reply.started":"2022-06-07T14:40:14.806908Z","shell.execute_reply":"2022-06-07T14:40:14.813249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# img = np.stack((R,\n#                G,\n#                B,\n#                Y), -1)\n# img = cv2.resize(img, (256,256))\n# img = np.divide(img, 255)\n# plt.imshow(img[:,:,:3])","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.815739Z","iopub.execute_input":"2022-06-07T14:40:14.816471Z","iopub.status.idle":"2022-06-07T14:40:14.822548Z","shell.execute_reply.started":"2022-06-07T14:40:14.816425Z","shell.execute_reply":"2022-06-07T14:40:14.821786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df['ID'].iloc[1], df['Label'].iloc[1]","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.824181Z","iopub.execute_input":"2022-06-07T14:40:14.824788Z","iopub.status.idle":"2022-06-07T14:40:14.833791Z","shell.execute_reply.started":"2022-06-07T14:40:14.824751Z","shell.execute_reply":"2022-06-07T14:40:14.833061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, df, path, size=None, label=True):        \n        self.df = df.copy()\n        self.path = path\n        self.size = size\n        self.label = label\n        if self.label:\n            self.df['Label'] = [[int(i) for i in s.split('|')] for s in self.df['Label']] # Data preprocessing\n            \n    def __getitem__(self, index):\n        img = open_rgby2(self.path, self.df['ID'].iloc[index], self.size)\n        if self.label:\n            target = np.eye(N_CLASSES,dtype=np.float)[self.df['Label'].iloc[index]].sum(axis=0) # One-hot encoding\n        else:\n            target = np.zeros(N_CLASSES,dtype=np.int)\n        return img, target\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.834763Z","iopub.execute_input":"2022-06-07T14:40:14.835544Z","iopub.status.idle":"2022-06-07T14:40:14.846061Z","shell.execute_reply.started":"2022-06-07T14:40:14.835506Z","shell.execute_reply":"2022-06-07T14:40:14.845357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"size = 256\nbatchSize = 32","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.847296Z","iopub.execute_input":"2022-06-07T14:40:14.847916Z","iopub.status.idle":"2022-06-07T14:40:14.85507Z","shell.execute_reply.started":"2022-06-07T14:40:14.847877Z","shell.execute_reply":"2022-06-07T14:40:14.854371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainLoader = DataLoader(CustomDataset(train_df, TRAIN, size), batch_size = batchSize, shuffle = True)\nvalidLoader = DataLoader(CustomDataset(valid_df, TRAIN, size), batch_size = batchSize, shuffle = True)\ntestLoader = DataLoader(CustomDataset(sub_df, TEST, size, False), batch_size=batchSize, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.857221Z","iopub.execute_input":"2022-06-07T14:40:14.857871Z","iopub.status.idle":"2022-06-07T14:40:14.898395Z","shell.execute_reply.started":"2022-06-07T14:40:14.857838Z","shell.execute_reply":"2022-06-07T14:40:14.897789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install efficientnet_pytorch\nfrom efficientnet_pytorch import EfficientNet","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:14.89954Z","iopub.execute_input":"2022-06-07T14:40:14.89979Z","iopub.status.idle":"2022-06-07T14:40:27.732343Z","shell.execute_reply.started":"2022-06-07T14:40:14.899759Z","shell.execute_reply":"2022-06-07T14:40:27.731546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# baseModel = EfficientNet.from_pretrained('efficientnet-b5', num_classes = 19)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:27.735888Z","iopub.execute_input":"2022-06-07T14:40:27.736102Z","iopub.status.idle":"2022-06-07T14:40:27.740525Z","shell.execute_reply.started":"2022-06-07T14:40:27.736074Z","shell.execute_reply":"2022-06-07T14:40:27.739719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# baseModel","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:27.742055Z","iopub.execute_input":"2022-06-07T14:40:27.744798Z","iopub.status.idle":"2022-06-07T14:40:27.749697Z","shell.execute_reply.started":"2022-06-07T14:40:27.74476Z","shell.execute_reply":"2022-06-07T14:40:27.748983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n\n    def forward(self, logit, target):\n        target = target.float()\n        max_val = (-logit).clamp(min=0)\n        loss = logit - logit * target + max_val + \\\n               ((-max_val).exp() + (-logit - max_val).exp()).log()\n\n        invprobs = F.logsigmoid(-logit * (target * 2.0 - 1.0))\n        loss = (invprobs * self.gamma).exp() * loss\n        if len(loss.size())==2:\n            loss = loss.sum(dim=1)\n        return loss.mean()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-06-07T14:40:27.75101Z","iopub.execute_input":"2022-06-07T14:40:27.751505Z","iopub.status.idle":"2022-06-07T14:40:27.760487Z","shell.execute_reply.started":"2022-06-07T14:40:27.751466Z","shell.execute_reply":"2022-06-07T14:40:27.759776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs('models')\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:27.761676Z","iopub.execute_input":"2022-06-07T14:40:27.762124Z","iopub.status.idle":"2022-06-07T14:40:27.836615Z","shell.execute_reply.started":"2022-06-07T14:40:27.762088Z","shell.execute_reply":"2022-06-07T14:40:27.835912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FinetunedEffiNetB5(nn.Module):\n    def __init__(self, n_classes = N_CLASSES):\n        super(FinetunedEffiNetB5, self).__init__()\n        baseModel = EfficientNet.from_pretrained('efficientnet-b5', num_classes = 19)\n#         baseModel._conv_stem = nn.Conv2d(4, 64, 7, 2, 3, bias=False)\n#         baseModel._conv_stem.in_channels = 4\n#         baseModel._conv_stem.weight = nn.Parameter(torch.cat([baseModel._conv_stem.weight, baseModel._conv_stem.weight], axis=1))\n#         baseModel._conv_stem = nn.Conv2d(4, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        self.efficientNet = baseModel\n        self.layer1 = nn.Linear(2000 , 256)\n        self.dropout = nn.Dropout(0.5)\n        self.layer2 = nn.Linear(256,n_classes)\n        self.ReLU = nn.LeakyReLU()\n        \n    def forward(self, input):\n        x  = self.efficientNet(input)\n        x = x.view(x.size(0),-1)\n        x = self.dropout(self.relu(self.layer1(x)))\n        x = self.layer2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:27.837977Z","iopub.execute_input":"2022-06-07T14:40:27.838572Z","iopub.status.idle":"2022-06-07T14:40:27.848433Z","shell.execute_reply.started":"2022-06-07T14:40:27.838535Z","shell.execute_reply":"2022-06-07T14:40:27.847479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testModel2 = EfficientNet.from_pretrained('efficientnet-b0', num_classes = 19)\ntestModel2","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:27.84994Z","iopub.execute_input":"2022-06-07T14:40:27.850454Z","iopub.status.idle":"2022-06-07T14:40:29.703014Z","shell.execute_reply.started":"2022-06-07T14:40:27.850416Z","shell.execute_reply":"2022-06-07T14:40:29.702195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# testModel = FinetunedEffiNetB5()\n# testModel.efficientNet._conv_stem.weight.size()","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:29.70447Z","iopub.execute_input":"2022-06-07T14:40:29.70476Z","iopub.status.idle":"2022-06-07T14:40:29.708198Z","shell.execute_reply.started":"2022-06-07T14:40:29.704723Z","shell.execute_reply":"2022-06-07T14:40:29.707261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# testModel.efficientNet","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:29.709475Z","iopub.execute_input":"2022-06-07T14:40:29.70996Z","iopub.status.idle":"2022-06-07T14:40:29.718901Z","shell.execute_reply.started":"2022-06-07T14:40:29.709925Z","shell.execute_reply":"2022-06-07T14:40:29.718077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.BCEWithLogitsLoss()\n# baseModel._conv_stem.in_channels = 4\n# baseModel._conv_stem.weight = torch.nn.Parameter(torch.cat([baseModel._conv_stem.weight, baseModel._conv_stem.weight], axis=1))\n# model = baseModel.to(device)\nmodel = testModel2\nmodel.cuda()\n\nlr = 0.001\noptimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=5e-4)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:29.721377Z","iopub.execute_input":"2022-06-07T14:40:29.721581Z","iopub.status.idle":"2022-06-07T14:40:32.867139Z","shell.execute_reply.started":"2022-06-07T14:40:29.721558Z","shell.execute_reply":"2022-06-07T14:40:32.866355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(epoch, history=None):\n    model.train()\n    t = tqdm(trainLoader)\n    \n    for batch_idx, (img_batch, label_batch) in enumerate(t):\n        img_batch = img_batch.to(device)\n        label_batch = label_batch.to(device)\n        \n        optimizer.zero_grad()\n        output = model(img_batch)\n        loss = criterion(output, label_batch)\n        t.set_description(f'train_loss (l={loss:.4f})')\n        \n        if history is not None:\n            history.loc[epoch + batch_idx / len(trainLoader), 'train_loss'] = loss.data.cpu().numpy()\n        \n        loss.backward()    \n        optimizer.step()\n    \n    torch.save(model.state_dict(), 'models/epoch{}.pth'.format(epoch))","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:32.86863Z","iopub.execute_input":"2022-06-07T14:40:32.868915Z","iopub.status.idle":"2022-06-07T14:40:32.876318Z","shell.execute_reply.started":"2022-06-07T14:40:32.86888Z","shell.execute_reply":"2022-06-07T14:40:32.875115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def binarize_prediction(probabilities, threshold: float, argsorted=None,\n                        min_labels=1, max_labels=10):\n    \"\"\" Return matrix of 0/1 predictions, same shape as probabilities.\n    \"\"\"\n    assert probabilities.shape[1] == N_CLASSES\n    if argsorted is None:\n        argsorted = probabilities.argsort(axis=1)\n    max_mask = _make_mask(argsorted, max_labels)\n    min_mask = _make_mask(argsorted, min_labels)\n    prob_mask = probabilities > threshold\n    return (max_mask & prob_mask) | min_mask\n\ndef _make_mask(argsorted, top_n: int):\n    mask = np.zeros_like(argsorted, dtype=np.uint8)\n    col_indices = argsorted[:, -top_n:].reshape(-1)\n    row_indices = [i // top_n for i in range(len(col_indices))]\n    mask[row_indices, col_indices] = 1\n    return mask","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:32.877817Z","iopub.execute_input":"2022-06-07T14:40:32.878076Z","iopub.status.idle":"2022-06-07T14:40:32.893016Z","shell.execute_reply.started":"2022-06-07T14:40:32.878042Z","shell.execute_reply":"2022-06-07T14:40:32.892294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eval(epoch, history =  None):\n    model.eval()\n    valid_loss = 0\n    all_predictions, all_targets = [], []\n    \n    with torch.no_grad():\n        for batch_idx, (img_batch, label_batch) in enumerate(validLoader):\n            all_targets.append(label_batch.numpy().copy())\n            img_batch = img_batch.to(device)\n            label_batch = label_batch.to(device)\n\n            output = model(img_batch)\n            loss = criterion(output, label_batch)\n            valid_loss += loss.data\n            predictions = torch.sigmoid(output)\n            all_predictions.append(predictions.cpu().numpy())\n    all_predictions = np.concatenate(all_predictions)\n    all_targets = np.concatenate(all_targets)\n    \n    valid_loss /= (batch_idx+1)\n    \n    if history is not None:\n        history.loc[epoch, 'valid_loss'] = valid_loss.cpu().numpy()\n    \n    print('Epoch: {}\\tLR: {:.6f}\\tValid Loss: {:.4f}'.format(\n        epoch, optimizer.state_dict()['param_groups'][0]['lr'], valid_loss))\n    \n    def get_score(y_pred):\n        with warnings.catch_warnings():\n            warnings.simplefilter('ignore', category=UndefinedMetricWarning)\n            return f1_score(all_targets, y_pred, average='macro')\n    \n    metrics = {}\n    argsorted = all_predictions.argsort(axis=1)\n    \n    return valid_loss","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:32.894301Z","iopub.execute_input":"2022-06-07T14:40:32.894829Z","iopub.status.idle":"2022-06-07T14:40:32.906421Z","shell.execute_reply.started":"2022-06-07T14:40:32.894794Z","shell.execute_reply":"2022-06-07T14:40:32.905674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history_train = pd.DataFrame()\nhistory_valid = pd.DataFrame()\n\nn_epochs = 3\ninit_epoch = 0\nmax_lr_changes = 1\nvalid_losses = []\nlr_reset_epoch = init_epoch\npatience = 2\nlr_changes = 0\nbest_valid_loss = 1000.\n\nfor epoch in range(init_epoch, n_epochs):\n    torch.cuda.empty_cache()\n    gc.collect()\n    train(epoch, history_train)\n    valid_loss = eval(epoch, history_valid)\n    valid_losses.append(valid_loss)\n\n    if valid_loss < best_valid_loss:\n        best_valid_loss = valid_loss\n    elif (patience and epoch - lr_reset_epoch > patience and\n          min(valid_losses[-patience:]) > best_valid_loss): \n        if lr_changes > max_lr_changes:\n            break\n        lr /= 5\n        print(f'lr updated to {lr}')\n        lr_reset_epoch = epoch\n        optimizer.param_groups[0]['lr'] = lr","metadata":{"execution":{"iopub.status.busy":"2022-06-07T14:40:32.909034Z","iopub.execute_input":"2022-06-07T14:40:32.910173Z","iopub.status.idle":"2022-06-07T18:30:10.631323Z","shell.execute_reply.started":"2022-06-07T14:40:32.910136Z","shell.execute_reply":"2022-06-07T18:30:10.630351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.memory_summary(device=None, abbreviated=False)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T20:10:18.887645Z","iopub.execute_input":"2022-06-07T20:10:18.88792Z","iopub.status.idle":"2022-06-07T20:10:18.953677Z","shell.execute_reply.started":"2022-06-07T20:10:18.887889Z","shell.execute_reply":"2022-06-07T20:10:18.951923Z"},"trusted":true},"execution_count":null,"outputs":[]}]}