{"cells":[{"metadata":{"_kg_hide-input":false,"trusted":true,"_uuid":"8ead08116baa59092b769e8f05de7426ef7c0faa","_kg_hide-output":true},"cell_type":"code","source":"!pip install fastai==0.7.0 --no-deps\n!pip install torch==0.4.1 torchvision==0.2.1","execution_count":null,"outputs":[]},{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","trusted":true},"cell_type":"code","source":"from fastai.conv_learner import *\nfrom fastai.dataset import *\nfrom fastai import *\n\nimport pandas as pd\nimport numpy as np\nimport os\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score\nimport scipy.optimize as opt\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"1dcb0ed47f29e7a457cc54500b97ab4e96405e8e"},"cell_type":"code","source":"PATH = './'\nTRAIN = '../input/human-protein-atlas-image-classification/train/'\nTEST1 = '../input/human-protein-atlas-image-classification/test/'\nLABELS = '../input/human-protein-atlas-image-classification/train.csv'\nSPLIT = '../input/protein-trainval-split/' \n# TEST='../input/pictures/';pic=TEST+'00008af0-bad0-11e8-b2b8-ac1f6b6435d0_blue.png'\nTEST='../input/picture-1/';pic=TEST+'0070171c-bad0-11e8-b2b8-ac1f6b6435d0_blue.png'\nnw = 2  ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"87ad005d93033f87ab77acf72d2d6f6cb4a33193"},"cell_type":"code","source":"name_label_dict = {\n0:  'Nucleoplasm',\n1:  'Nuclear membrane',\n2:  'Nucleoli',   \n3:  'Nucleoli fibrillar center',\n4:  'Nuclear speckles',\n5:  'Nuclear bodies',\n6:  'Endoplasmic reticulum',   \n7:  'Golgi apparatus',\n8:  'Peroxisomes',\n9:  'Endosomes',\n10:  'Lysosomes',\n11:  'Intermediate filaments',\n12:  'Actin filaments',\n13:  'Focal adhesion sites',   \n14:  'Microtubules',\n15:  'Microtubule ends',  \n16:  'Cytokinetic bridge',   \n17:  'Mitotic spindle',\n18:  'Microtubule organizing center',  \n19:  'Centrosome',\n20:  'Lipid droplets',\n21:  'Plasma membrane',   \n22:  'Cell junctions', \n23:  'Mitochondria',\n24:  'Aggresome',\n25:  'Cytosol',\n26:  'Cytoplasmic bodies',   \n27:  'Rods & rings' }","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"d2376dd252c2cd9a88fd60866e0ac4bec8ea1314"},"cell_type":"code","source":"with open(os.path.join(SPLIT,'tr_names.txt'), 'r') as text_file:\n    tr_n = text_file.read().split(',')\nwith open(os.path.join(SPLIT,'val_names.txt'), 'r') as text_file:\n    val_n = text_file.read().split(',')\ntest_names = sorted({f[:36] for f in os.listdir(TEST)})\n# pic_name=sorted({f[:36] for f in os.listdir(PICTURE)})\n# print(pic_name)\nprint(len(tr_n),len(val_n))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"e14c49dde171f5ffb21cbe67a7d4bb68a88e5084"},"cell_type":"code","source":"class Oversampling:\n    def __init__(self,path):\n        self.train_labels = pd.read_csv(path).set_index('Id')\n        self.train_labels['Target'] = [[int(i) for i in s.split()] \n                                       for s in self.train_labels['Target']]  \n        #set the minimum number of duplicates for each class\n        self.multi = [1,1,1,1,1,1,1,1,\n                      4,4,4,1,1,1,1,4,\n                      1,1,1,1,2,1,1,1,\n                      1,1,1,4]\n\n    def get(self,image_id):\n        labels = self.train_labels.loc[image_id,'Target'] if image_id \\\n          in self.train_labels.index else []\n        m = 1\n        for l in labels:\n            if m < self.multi[l]: m = self.multi[l]\n        return m\n    \ns = Oversampling(os.path.join(PATH,LABELS))\ntr_n = [idx for idx in tr_n for _ in range(s.get(idx))]\nprint(len(tr_n),flush=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"b5ae50160b85e023518ee17d62e4ff216eb5b9fb"},"cell_type":"code","source":"def open_rgby(path,id): #a function that reads RGBY image\n    colors = ['red','green','blue','yellow']\n    flags = cv2.IMREAD_GRAYSCALE\n    img = [cv2.imread(os.path.join(path, id+'_'+color+'.png'), flags).astype(np.float32)/255\n           for color in colors]\n    return np.stack(img, axis=-1)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"0dd72024e421ac2ced8a21389e15e914fae8d6a8"},"cell_type":"code","source":"class pdFilesDataset(FilesDataset):\n    def __init__(self, fnames, path, transform):\n        self.labels = pd.read_csv(LABELS).set_index('Id')\n        self.labels['Target'] = [[int(i) for i in s.split()] for s in self.labels['Target']]\n        super().__init__(fnames, transform, path)\n    \n    def get_x(self, i):\n        return open_rgby(self.path,self.fnames[i])\n    \n    def get_y(self, i):\n        if(self.path == TEST): return np.zeros(len(name_label_dict),dtype=np.int)\n        else:\n            labels = self.labels.loc[self.fnames[i]]['Target']\n            return np.eye(len(name_label_dict),dtype=np.float)[labels].sum(axis=0)\n        \n    @property\n    def is_multi(self): return True\n    @property\n    def is_reg(self):return True\n    #this flag is set to remove the output sigmoid that allows log(sigmoid) optimization\n    #of the numerical stability of the loss function\n    \n    def get_c(self): return len(name_label_dict) #number of classes","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"2279cbaea32d80e9bf030b4a3e74ade811a52b3a"},"cell_type":"code","source":"def get_data(sz,bs,is_test=False):\n    #data augmentation\n    if is_test:\n        aug_tfms = [RandomRotate(30, tfm_y=TfmType.NO),\n                RandomDihedral(tfm_y=TfmType.NO)]\n    else:\n        aug_tfms = [RandomRotate(30, tfm_y=TfmType.NO),\n                RandomDihedral(tfm_y=TfmType.NO),\n                RandomLighting(0.05, 0.05, tfm_y=TfmType.NO),\n                Cutout(n_holes=25, length=10*sz//128, tfm_y=TfmType.NO)]\n    #mean and std in of each channel in the train set\n    stats = A([0.08069, 0.05258, 0.05487, 0.08282], [0.13704, 0.10145, 0.15313, 0.13814])\n    tfms = tfms_from_stats(stats, sz, crop_type=CropType.NO, tfm_y=TfmType.NO, \n                aug_tfms=aug_tfms)\n    ds = ImageData.get_ds(pdFilesDataset, (tr_n[:-(len(tr_n)%bs)],TRAIN), \n                (val_n,TRAIN), tfms, test=(test_names,TEST))\n    md = ImageData(PATH, ds, bs, num_workers=nw, classes=None)\n    return md","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"24adecf3a1b558b05026ec43e64b8848ac0b05d7"},"cell_type":"code","source":"bs = 16\nsz = 256\nmd = get_data(sz,bs,is_test=True)\n\nx,y = next(iter(md.aug_dl))\nx.shape, y.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"c288311b2416d46e5ec0126c92ccfed170ded223"},"cell_type":"code","source":"def display_imgs(x):\n    columns = 4\n    bs = x.shape[0]\n    rows = min((bs+3)//4,4)\n    fig=plt.figure(figsize=(columns*4, rows*4))\n    for i in range(rows):\n        for j in range(columns):\n            idx = i+j*columns\n            fig.add_subplot(rows, columns, idx+1)\n            plt.axis('off')\n            plt.imshow((x[idx,:,:,:3]*255).astype(np.int))\n    plt.show()\n    \n# display_imgs(np.asarray(md.trn_ds.denorm(x)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"4f82ac4cd3a7aaf5b61a5ffa4e882ee0dcec125c"},"cell_type":"code","source":"class Resnet34_4(nn.Module):\n    def __init__(self, pre=True):\n        super().__init__()\n        encoder = resnet34(pretrained=pre)\n        \n        self.conv1 = nn.Conv2d(4, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        if(pre):\n            w = encoder.conv1.weight\n            self.conv1.weight = nn.Parameter(torch.cat((w,\n                                    0.5*(w[:,:1,:,:]+w[:,2:,:,:])),dim=1))\n        \n        self.bn1 = encoder.bn1\n        self.relu = nn.ReLU(inplace=True) \n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        self.layer0 = nn.Sequential(self.conv1,self.relu,self.bn1,self.maxpool)\n        self.layer1 = encoder.layer1\n        self.layer2 = encoder.layer2\n        self.layer3 = encoder.layer3\n        self.layer4 = encoder.layer4\n        #the head will be added automatically by fast.ai\n        \n    def forward(self, x):\n        x = self.layer0(x)\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        \n        return x","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"b3f8762045c4087ae9aecd22f8c13a075f567d7a"},"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, 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        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()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"06c8f2af90481b71e1f174c584b7d1d2a0697363"},"cell_type":"code","source":"def acc(preds,targs,th=0.0):\n    preds = (preds > th).int()\n    targs = targs.int()\n    return (preds==targs).float().mean()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"30197d3107d0dea338fb88c1b5c50ec356f9e17d"},"cell_type":"code","source":"class F1:\n    __name__ = 'F1 macro'\n    def __init__(self,n=28):\n        self.n = n\n        self.TP = np.zeros(self.n)\n        self.FP = np.zeros(self.n)\n        self.FN = np.zeros(self.n)\n\n    def __call__(self,preds,targs,th=0.0):\n        preds = (preds > th).int()\n        targs = targs.int()\n        self.TP += (preds*targs).float().sum(dim=0)\n        self.FP += (preds > targs).float().sum(dim=0)\n        self.FN += (preds < targs).float().sum(dim=0)\n        score = (2.0*self.TP/(2.0*self.TP + self.FP + self.FN + 1e-6)).mean()\n        return score\n\n    def reset(self):\n        #macro F1 score\n        score = (2.0*self.TP/(2.0*self.TP + self.FP + self.FN + 1e-6))\n        print('F1 macro:',score.mean(),flush=True)\n        #print('F1:',score)\n        self.TP = np.zeros(self.n)\n        self.FP = np.zeros(self.n)\n        self.FN = np.zeros(self.n)\n\nclass F1_callback(Callback):\n    def __init__(self, n=28):\n        self.f1 = F1(n)\n\n    def on_epoch_end(self, metrics):\n        self.f1.reset()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"17c6374e2043d19b460ae425eb4009855423bc27","scrolled":true,"_kg_hide-output":true},"cell_type":"code","source":"sz = 256 #image size\nbs = 64  #batch size\n\nmd = get_data(sz,bs)\nlearner = ConvLearner.pretrained(Resnet34_4, md, ps=0.5) #dropout 50%\nlearner.opt_fn = optim.Adam\nlearner.clip = 1.0 #gradient clipping\nlearner.crit = FocalLoss()\nf1_callback = F1_callback()\nlearner.metrics = [acc,f1_callback.f1]\nlearner.summary","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!cp ../input/resnet34-256-1h5/ResNet34_256_1.h5 models/resnet34_256_1.h5\n!cp ../input/pictures/00008af0-bad0-11e8-b2b8-ac1f6b6435d0_red.png test.png","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learner.load(\"resnet34_256_1\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"61cb0f0ceb2141d4bd3d2f2b181a207f7b503dd2"},"cell_type":"code","source":"# with warnings.catch_warnings():\n#     warnings.simplefilter(\"ignore\")\n#     learner.lr_find()\n# learner.sched.plot()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"fa183a7b7809e4cbd6aa16916d9915a24826d58b"},"cell_type":"code","source":"# lr = 0.5e-2\n# with warnings.catch_warnings():\n#     warnings.simplefilter(\"ignore\")\n#     learner.fit(lr,1,callbacks=[f1_callback])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"032d2c930d002704e5690f095fc0af043c739982"},"cell_type":"code","source":"# learner.unfreeze()\n# lrs=np.array([lr/10,lr/3,lr])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"7b01fe75583db3453bc3712e4c7d9e36f2e5bdb0"},"cell_type":"code","source":"# with warnings.catch_warnings():\n#     warnings.simplefilter(\"ignore\")\n#     learner.fit(lrs/4,4,cycle_len=2,use_clr=(10,20),callbacks=[f1_callback])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"f36c0b2468425ee0f064decbc1ee40777c84eb52"},"cell_type":"code","source":"# with warnings.catch_warnings():\n#     warnings.simplefilter(\"ignore\")\n#     learner.fit(lrs/16,2,cycle_len=4,use_clr=(10,20),callbacks=[f1_callback])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"d20a5444b01a8ab1773ce0beb7e253f90d1ffe72"},"cell_type":"code","source":"# with warnings.catch_warnings():\n#     warnings.simplefilter(\"ignore\")\n#     learner.fit(lrs/32,1,cycle_len=8,use_clr=(10,20),callbacks=[f1_callback])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"032a66e6f9db41870dbbbfa5ce7a106363f8b8e0"},"cell_type":"code","source":"# learner.save('ResNet34_256_1')","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"83132c40a82e38beb33ba62fa19e5a0233a9cbbb"},"cell_type":"markdown","source":"### Submission"},{"metadata":{"trusted":true,"_uuid":"cedaa4bf73c9b3381ac3e6da32b040e05ce24b91","scrolled":true},"cell_type":"code","source":"preds_t,y_t = learner.TTA(n_aug=8,is_test=True)\npreds_t = np.stack(preds_t, axis=-1)\npred_t = preds_t.mean(axis=-1)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom PIL import Image\n\nimg=Image.open(pic)\nplt.imshow(img)\nplt.axis('off')\nplt.show()\nfor line in pred_t:\n    results = ' '.join(list([str(i) for i in np.nonzero(line>0)[0]]))\n    print('预测结果为：',results)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def save_pred(pred, th=0.0, fname='protein_classification.csv'):\n    pred_list = []\n    for line in pred:\n        s = ' '.join(list([str(i) for i in np.nonzero(line>th)[0]]))\n        pred_list.append(s)\n    \n    df = pd.DataFrame({'Id':learner.data.test_ds.fnames,'Predicted':pred_list})\n    df.sort_values(by='Id').to_csv(fname, header=True, index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"save_pred(pred_t,0,'protein_classification.csv')","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}