{"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":"# prepare to use TPU\n\n- install torch-xla\n- import packages","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y torch torchvision fastai easyocr allennlp torchtext torchaudio pytorch-lightning kornia fairscale\n!curl https://raw.githubusercontent.com/pytorch/xla/4e3de8c52323acd62c87ad3107731c657c9860e7/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n!python pytorch-xla-env-setup.py\n!pip install ../input/torchmetrics/torchmetrics-0.9.1-py3-none-any.whl\n!rm *.whl\n!rm pytorch-xla-env-setup.py\n!pip install torch_optimizer","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-02-24T08:00:43.446545Z","iopub.execute_input":"2023-02-24T08:00:43.446965Z","iopub.status.idle":"2023-02-24T08:01:49.045394Z","shell.execute_reply.started":"2023-02-24T08:00:43.446854Z","shell.execute_reply":"2023-02-24T08:01:49.044126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport sys\nimport time\nimport h5py\nimport os\nimport gc\nimport cv2\nimport math\nimport random\nimport pickle\nimport pydicom as dicom\nfrom PIL import Image\nimport torch\nfrom torch import nn, optim\nimport torchvision\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import StepLR\nsys.path.append(\"../input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master\")\nfrom efficientnet_pytorch import model as enet\n\nfrom tqdm.notebook import tqdm\nfrom sklearn.model_selection import StratifiedKFold\nfrom multiprocessing import Pool\n\n# 乱数を初期化する\nrandom.seed(42)\nnp.random.seed(42)\ntorch.manual_seed(42)","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:49.048385Z","iopub.execute_input":"2023-02-24T08:01:49.048852Z","iopub.status.idle":"2023-02-24T08:01:50.359307Z","shell.execute_reply.started":"2023-02-24T08:01:49.048795Z","shell.execute_reply":"2023-02-24T08:01:50.358386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metadata","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 8\nBATCH_SIZE_VAL = 2\nNUM_EPOCHS = 6\nUPDATE_DA_EPOCH = 0,2,4\nGRAD_CLIP = True\nIMAGE_SIZE = 512\nNUM_TRAIN = 40000\nNUM_TEST = 240\nNUM_VALID = 1000\nTEST_RUN = False\nUSE_MODEL = ('efficientnet-b6','../input/efficientnet-pytorch/efficientnet-b6-c76e70fd.pth',2308)","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:50.360989Z","iopub.execute_input":"2023-02-24T08:01:50.361937Z","iopub.status.idle":"2023-02-24T08:01:50.368246Z","shell.execute_reply.started":"2023-02-24T08:01:50.361882Z","shell.execute_reply":"2023-02-24T08:01:50.367076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TEST_RUN:\n    NUM_TRAIN = 220\n    NUM_TEST = 50\n    NUM_VALID = 50\n    NUM_EPOCHS = 3\n    UPDATE_DA_EPOCH = 0,1,2\n    USE_MODEL = ('efficientnet-b1','/kaggle/input/efficientnet-pytorch/efficientnet-b1-dbc7070a.pth',1284)","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:50.371319Z","iopub.execute_input":"2023-02-24T08:01:50.371785Z","iopub.status.idle":"2023-02-24T08:01:50.386535Z","shell.execute_reply.started":"2023-02-24T08:01:50.371742Z","shell.execute_reply":"2023-02-24T08:01:50.38571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make validation data\n\n**score is caluclate within \"CC\" and \"MLO\" views, so data split with view**\n\n```\npatient1-L - CC   ---+ \npatient1-L - MLO  ---+----> predicted value to validate\n```\n\n**if:**\n\n```\npatient1-L - CC = 0\npatient1-L - MLO = 1     ----> predicted value to valudate should 1\n```\n\n```\npatient1-L - CC = 1\npatient1-L - MLO = 0     ----> predicted value to valudate  should 1\n```\n\n```\npatient1-L - CC = 0\npatient1-L - MLO = 0     ----> predicted value to valudate  should 0\n```","metadata":{}},{"cell_type":"code","source":"# Train metadata\ndef get_datas():\n    df = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\n    df = df.drop(df[(df.view==\"AT\")|(df.view==\"LM\")|(df.view==\"ML\")|(df.view==\"LMO\")].index)\n    df = df.reset_index()\n    df[\"type\"] = \"train\"\n\n    # Make validation data\n    #  score is caluclate within \"CC\" and \"MLO\" views, so data split with view\n    dt = df[49997:] # data to validation\n    dtc = dt[dt[\"view\"]==\"CC\"]\n    dtm = dt[dt[\"view\"]==\"MLO\"]\n    dfc_testtrue = dtc[dtc.cancer==1]\n    dfm_testtrue = dtm[dtm.cancer==1]\n\n    cancers = set(dfc_testtrue.patient_id.values.tolist()+dfm_testtrue.patient_id.values.tolist())\n    nocancers = sorted(list(set(dt.patient_id.values.tolist())-cancers)) # who have no cancer\n    cancers = sorted(list(cancers)) # who have cancer\n\n    # split a haft to test and validation\n    cancers_test = cancers[:len(cancers)//2]\n    cancers_valid = cancers[len(cancers)//2:]\n    nocancers_test = nocancers[:len(nocancers)//2]\n    npcancers_valid = nocancers[len(nocancers)//2:]\n\n    # Make test/validation data for cancer and no-cancer\n    df_cancers_test = pd.concat([dt[dt.patient_id == c] for c in cancers_test], axis=0)\n    df_cancers_valid = pd.concat([dt[dt.patient_id == c] for c in cancers_valid], axis=0)\n    df_nocancers_test = pd.concat([dt[dt.patient_id == c] for c in nocancers_test], axis=0)\n    df_nocancers_valid = pd.concat([dt[dt.patient_id == c] for c in npcancers_valid], axis=0)\n    \n    df_test = pd.concat([df_cancers_test[:NUM_TEST//2], df_nocancers_test[:NUM_TEST//2]], axis=0)\n    df_valid = pd.concat([df_cancers_valid[:NUM_VALID//2], df_nocancers_valid[:NUM_VALID//2]], axis=0)\n    \n    # Drop patients unless 2view and 2laterality has\n    def droppaient(df):\n        drop_test = []\n        for i,p in zip(df.index,df.patient_id):\n            l = [len(df[(df.patient_id==p)&(df[\"view\"]==\"CC\")&(df.laterality==\"L\")]),\n                 len(df[(df.patient_id==p)&(df[\"view\"]==\"CC\")&(df.laterality==\"R\")]),\n                 len(df[(df.patient_id==p)&(df[\"view\"]==\"MLO\")&(df.laterality==\"L\")]),\n                 len(df[(df.patient_id==p)&(df[\"view\"]==\"MLO\")&(df.laterality==\"R\")])]\n            if np.min(l) == 0:\n                drop_test.append(i)\n        return df.drop(index=drop_test)\n    \n    df_test = droppaient(df_test)\n    df_valid = droppaient(df_valid)\n\n    # Training data\n    dt = df[:49997]\n    df_train = pd.concat([dt[dt.cancer==1][:NUM_TRAIN//2], dt[dt.cancer==0][:NUM_TRAIN//2]], axis=0)\n\n    # reset index\n    df_train = df_train.reset_index()\n    df_test = df_test.reset_index()\n    df_valid = df_valid.reset_index()\n    \n    return df_train, df_test, df_valid\n\ndf_train, df_test, df_valid = get_datas()\ngc.collect()\n\npd.DataFrame({\"DataLength\":[len(df_train), len(df_test), len(df_valid)], \"CancerRate\":[np.mean(df_train.cancer.values), np.mean(df_test.cancer.values), np.mean(df_valid.cancer.values)]}, index=[\"train\",\"test\",\"valid\"])","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:50.388228Z","iopub.execute_input":"2023-02-24T08:01:50.388551Z","iopub.status.idle":"2023-02-24T08:01:52.165537Z","shell.execute_reply.started":"2023-02-24T08:01:50.388513Z","shell.execute_reply":"2023-02-24T08:01:52.164433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get Uniqed index for each patients by view and laterality","metadata":{}},{"cell_type":"code","source":"def get_first_index(df, patient, view, late):\n    return df[(df.patient_id==patient)&(df[\"view\"]==view)&(df.laterality==late)].index[0]\ntest_patients = sorted(list(set(df_test.patient_id.values.tolist())))\ntest_index_CCL = [get_first_index(df_test, p, \"CC\", \"L\") for p in test_patients]\ntest_index_CCR = [get_first_index(df_test, p, \"CC\", \"R\") for p in test_patients]\ntest_index_MLOL = [get_first_index(df_test, p, \"MLO\", \"L\") for p in test_patients]\ntest_index_MLOR = [get_first_index(df_test, p, \"MLO\", \"R\") for p in test_patients]\nvalid_patients = sorted(list(set(df_valid.patient_id.values.tolist())))\nvalid_index_CCL = [get_first_index(df_valid, p, \"CC\", \"L\") for p in valid_patients]\nvalid_index_CCR = [get_first_index(df_valid, p, \"CC\", \"R\") for p in valid_patients]\nvalid_index_MLOL = [get_first_index(df_valid, p, \"MLO\", \"L\") for p in valid_patients]\nvalid_index_MLOR = [get_first_index(df_valid, p, \"MLO\", \"R\") for p in valid_patients]\ndel get_first_index, test_patients, valid_patients\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:52.167007Z","iopub.execute_input":"2023-02-24T08:01:52.167332Z","iopub.status.idle":"2023-02-24T08:01:52.394866Z","shell.execute_reply.started":"2023-02-24T08:01:52.1673Z","shell.execute_reply":"2023-02-24T08:01:52.39395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Parameters","metadata":{}},{"cell_type":"code","source":"import torch_optimizer\n\nclass FocalLoss(nn.Module):\n    \"\"\"\n    The focal loss for fighting against class-imbalance\n    \"\"\"\n\n    def __init__(self, alpha=1, gamma=2):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.epsilon = 1e-12  # prevent training from Nan-loss error\n\n    def forward(self, probs, target):\n        \"\"\"\n        logits & target should be tensors with shape [batch_size, num_classes]\n        \"\"\"\n        #probs = F.sigmoid(logits)\n        one_subtract_probs = 1.0 - probs\n        # add epsilon\n        probs_new = probs + self.epsilon\n        one_subtract_probs_new = one_subtract_probs + self.epsilon\n        # calculate focal loss\n        log_pt = target * torch.log(probs_new) + (1.0 - target) * torch.log(one_subtract_probs_new)\n        pt = torch.exp(log_pt)\n        focal_loss = -1.0 * (self.alpha * (1 - pt) ** self.gamma) * log_pt\n        return torch.mean(focal_loss)\n\ndef get_OPTIMIZER(p, lr):\n    opt = torch_optimizer.QHM(\n                p,\n                lr=lr,\n                momentum=0.99,\n                nu=0.7,\n                weight_decay=5e-6,\n                weight_decay_type='grad',\n            )\n    return opt\n    #return torch.optim.Adam(p, lr=lr)\n\ndef get_LOSS():\n    return FocalLoss()","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:52.396615Z","iopub.execute_input":"2023-02-24T08:01:52.397288Z","iopub.status.idle":"2023-02-24T08:01:52.417801Z","shell.execute_reply.started":"2023-02-24T08:01:52.397233Z","shell.execute_reply":"2023-02-24T08:01:52.416656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# use rsna-mammography-images-as-pngs thank you Radek\n\n#def one_da(i):\n#    r = df.iloc[i].astype(str)\n#    image_id = r.image_id\n#    patient_id = r.patient_id\n#    data_type = r.type\n#    dirname = '/kaggle/input/rsna-breast-cancer-detection/%s_images/%s/' % (data_type, patient_id)\n#    fn = image_id+\".dcm\"\n#    try:\n#        ds = dicom.dcmread(dirname+fn)\n#        img = ds.pixel_array\n#        img = (img - img.min()) / (img.max() - img.min())\n#        if ds.PhotometricInterpretation == \"MONOCHROME1\":\n#            img = 1 - img\n#        img = cv2.resize(img, (IMAGE_SIZE,IMAGE_SIZE))\n#        img = (255*img).astype(np.uint8)\n#        cv2.imwrite(\"../temp/%s-%s.jpg\"%(data_type,image_id), img)\n#    except:\n#        pass\n#\n#def da():\n#    with Pool(4) as pool:\n#        with tqdm(total=len(df)) as t:\n#            for _ in pool.imap_unordered(one_da, list(range(len(df)))):\n#                t.update(1)\n#\n#da()\n#gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:52.419329Z","iopub.execute_input":"2023-02-24T08:01:52.419596Z","iopub.status.idle":"2023-02-24T08:01:52.424671Z","shell.execute_reply.started":"2023-02-24T08:01:52.419566Z","shell.execute_reply":"2023-02-24T08:01:52.423701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Setting","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass SiLU(nn.Module):\n    def forward(self,x):\n        return x * F.sigmoid(x)\n    \nclass Attention(nn.Module):\n    def __init__(self, channel, head_ch=64):\n        super().__init__()\n        self.n_heads = channel // head_ch\n\n        self.q = nn.Linear(channel, channel, bias=False)\n        self.k = nn.Linear(channel, channel, bias=False)\n        self.v = nn.Linear(channel, channel, bias=False)\n\n        self.out = nn.Linear(channel, channel)\n\n    def forward(self, x):\n        bs, length, channel = x.shape\n        ch = channel // self.n_heads\n\n        q = self.q(x) # b,l,c\n        k = self.k(x) # b,l,c\n        v = self.v(x) # b,l,c\n        q = torch.reshape(q, (bs, length, self.n_heads, ch))\n        k = torch.reshape(k, (bs, length, self.n_heads, ch))\n        v = torch.reshape(v, (bs, length, self.n_heads, ch))\n        q = torch.transpose(q, 1, 2) # b,h,l,c\n        k = torch.transpose(k, 1, 2) # b,h,l,c\n        v = torch.transpose(v, 1, 2) # b,h,l,c\n        q = torch.reshape(q, (bs * self.n_heads, length, ch)) # bh,l,c\n        k = torch.reshape(k, (bs * self.n_heads, length, ch)) # bh,l,c\n        v = torch.reshape(v, (bs * self.n_heads, length, ch)) # bh,l,c\n\n        k = torch.transpose(k, 1, 2)   # bh,c,l\n        weight = torch.bmm(q,k)    # bh,l,l\n        weight = weight * (int(ch)**(-0.5))\n        weight = torch.nn.functional.softmax(weight, dim=-1)\n\n        # attend to values\n        v = torch.transpose(v, 1, 2)   # bh,c,l\n        hidden = torch.bmm(v,weight)     # bh,c,l\n        hidden = torch.reshape(hidden, (bs, self.n_heads, length, ch)) # b,h,l,c\n        hidden = torch.transpose(hidden, 1, 2) # b,l,h,c\n        hidden = torch.reshape(hidden, (bs, length, channel)) # b,l,c\n\n        return self.out(hidden)\n\nclass Transformer(nn.Module):\n    def __init__(self, channel, channel_head, mlp_channel, dropout=0.):\n        super().__init__()\n        self.att = nn.Sequential(\n            nn.LayerNorm(channel),\n            Attention(channel, channel_head),\n        )\n        self.ffn = nn.Sequential(\n            nn.LayerNorm(channel),\n            nn.Linear(channel, mlp_channel),\n            SiLU(),\n            nn.Linear(mlp_channel, channel),\n            nn.Dropout(dropout),\n        )\n\n    def forward(self, h):\n        b, d, x, y = h.shape\n        h = torch.reshape(h, (b, d, x*y))\n        h = torch.transpose(h, 1, 2)\n        h = self.att(h) + h\n        h = self.ffn(h) + h\n        h = torch.transpose(h, 1, 2)\n        h = torch.reshape(h, (b, d, x, y))\n        return h\n\nclass PositionalEmbedding(nn.Module):\n    def __init__(self, channel, block=4):\n        super().__init__()\n        inv_freq = 1 / (10000 ** (torch.arange(0.0, channel, 2.0) / channel))\n        self.register_buffer(\"inv_freq\", inv_freq)\n        self.ff = nn.Linear(channel//block, channel//block)\n        self.channel = channel\n        self.block = block\n\n    def forward(self, h):\n        dtype = h.dtype\n        b, d, x, y = h.shape\n        h = torch.reshape(h, (b, d, x*y))\n        h = torch.transpose(h, 1, 2)\n\n        pos_seq = torch.arange(x*y).type(dtype).to(h.device)\n        sinusoid_inp = torch.ger(pos_seq, self.inv_freq)\n        pos_emb = torch.cat([sinusoid_inp.sin(), sinusoid_inp.cos()], dim=-1)\n        pos_emb = torch.reshape(pos_emb, (x*y, self.block, self.channel//self.block))\n        pos_emb = self.ff(pos_emb)\n        pos_emb = torch.reshape(pos_emb, (x*y, self.channel))\n\n        h = h + pos_emb.unsqueeze(0)\n\n        h = torch.transpose(h, 1, 2)\n        h = torch.reshape(h, (b, d, x, y))\n        return h\n\nclass SeBlock(nn.Module):\n    def __init__(self,channel,stride=1, r=8):\n        super(SeBlock, self).__init__()\n        self.conv = nn.Conv2d(channel,channel,3,1,1,groups=channel)\n        self.ffn = nn.Sequential(\n            nn.Conv2d(channel,channel//8,1,1,bias=False),\n            SiLU(),\n            nn.Conv2d(channel//8,channel,1,1,bias=False),\n            nn.Sigmoid(),\n        )\n        self.norm = nn.BatchNorm2d(channel)\n\n    def forward(self,inp):\n        x = self.conv(inp)\n        h = F.adaptive_avg_pool2d(x, (1, 1))\n        x *= self.ffn(h)\n        return self.norm(x)\n","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:52.426387Z","iopub.execute_input":"2023-02-24T08:01:52.42674Z","iopub.status.idle":"2023-02-24T08:01:52.456886Z","shell.execute_reply.started":"2023-02-24T08:01:52.426703Z","shell.execute_reply":"2023-02-24T08:01:52.455834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class WDNet(nn.Module):\n    def __init__(self, n_class=1, n_depth=[2,4,8,8,6], is_vit=[0,0,0,1,1], n_channels=[64,128,256,384,512]):\n        super(WDNet, self).__init__()\n        channel = 42\n        layers = [nn.Conv2d(3, channel, 3, 1, 1),\n                        SiLU(),\n                        nn.Conv2d(channel,channel,3,1,1,groups=channel),\n                        SiLU(),\n                        nn.BatchNorm2d(channel)]\n        for i in range(len(n_depth)):\n            in_channel = channel\n            channel = n_channels[i]\n            layers.extend([\n                nn.Conv2d(in_channel,channel,1,1,bias=False),\n                SiLU(),\n                nn.Conv2d(channel,channel,3,1,1,groups=channel,bias=False),\n                SiLU(),\n                nn.BatchNorm2d(channel)\n            ])\n            if is_vit[i] == 0:\n                for _ in range(n_depth[i]-1):\n                    layers.append(SeBlock(channel))\n            else:\n                layers.append(PositionalEmbedding(channel))\n                for _ in range(n_depth[i]-1):\n                    head_ch = channel // (channel // 64)\n                    layers.append(Transformer(channel, head_ch, channel*2))\n            if i < len(n_depth)-1:\n                layers.extend([\n                    nn.Conv2d(channel,channel,2,2,bias=False),\n                ])\n\n        self.layer = nn.Sequential(*layers)\n        self.fc = nn.Linear(channel,n_class)\n\n    def forward(self, input):\n        x = self.layer(input)\n        x = F.adaptive_avg_pool2d(x, (1, 1))\n        x = x.view(x.size(0), -1)\n        result = self.fc(x)\n        return torch.sigmoid(result)\n","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:52.460014Z","iopub.execute_input":"2023-02-24T08:01:52.460345Z","iopub.status.idle":"2023-02-24T08:01:52.476028Z","shell.execute_reply.started":"2023-02-24T08:01:52.4603Z","shell.execute_reply":"2023-02-24T08:01:52.475046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TPU Setting","metadata":{}},{"cell_type":"code","source":"import torch_xla\nimport torch_xla.core.xla_model as xm\nfrom torch.utils.data.distributed import DistributedSampler\nimport torch_xla.distributed.parallel_loader as pl                             \nimport torch_xla.distributed.xla_multiprocessing as xmp\n#from torchmetrics import AUROC\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:52.477745Z","iopub.execute_input":"2023-02-24T08:01:52.478049Z","iopub.status.idle":"2023-02-24T08:01:52.682882Z","shell.execute_reply.started":"2023-02-24T08:01:52.477995Z","shell.execute_reply":"2023-02-24T08:01:52.681984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xm.get_xla_supported_devices()","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:52.684305Z","iopub.execute_input":"2023-02-24T08:01:52.684588Z","iopub.status.idle":"2023-02-24T08:01:58.053842Z","shell.execute_reply.started":"2023-02-24T08:01:52.684554Z","shell.execute_reply":"2023-02-24T08:01:58.052781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xm.xrt_world_size()","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.05506Z","iopub.execute_input":"2023-02-24T08:01:58.055855Z","iopub.status.idle":"2023-02-24T08:01:58.064769Z","shell.execute_reply.started":"2023-02-24T08:01:58.055814Z","shell.execute_reply":"2023-02-24T08:01:58.064035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Augumentation Method\n\n**Make Strengthen DA in 3 stages**\n\n1. flips, rotation, shift\n2. 1 and cutting paste\n3. 2 and mixed in","metadata":{}},{"cell_type":"code","source":"# Data Augumentation\n\ndef random_paste(img, img2):\n    img_height, img_width, _ = img.shape\n    patch_height, patch_width, _ = img2.shape\n    width_half = round(img_width / 2)\n    height_half = round(img_height / 2)\n    paste = cv2.resize(img2, (width_half, height_half))\n    r = np.random.randint(4)\n    if r == 0:\n        img[:height_half,:width_half,:] = paste\n    elif r == 1:\n        img[-height_half:,-width_half:,:] = paste\n    elif r == 2:\n        img[-height_half:,:width_half,:] = paste\n    elif r == 3:\n        img[:height_half,-width_half:,:] = paste\n    return img\n\ndef fill(img, h, w):\n    img = cv2.resize(img, (h, w), cv2.INTER_CUBIC)\n    return img\n        \ndef horizontal_shift(img, ratio=0.4):\n    if ratio > 1 or ratio < 0:\n        print('Value should be less than 1 and greater than 0')\n        return img\n    ratio = random.uniform(-ratio, ratio)\n    h, w = img.shape[:2]\n    to_shift = w*ratio\n    if ratio > 0:\n        img = img[:, :int(w-to_shift), :]\n    if ratio < 0:\n        img = img[:, int(-1*to_shift):, :]\n    img = fill(img, h, w)\n    return img\n\ndef random_rotation(img, angle=12):\n    angle = int(random.uniform(-angle, angle))\n    h, w = img.shape[:2]\n    M = cv2.getRotationMatrix2D((int(w/2), int(h/2)), angle, 1)\n    img = cv2.warpAffine(img, M, (w, h))\n    return img\n\ndef random_flip(img):\n    v = np.random.randint(4)\n    if v==3:\n        return img\n    return cv2.flip(img, v-1)","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.06592Z","iopub.execute_input":"2023-02-24T08:01:58.066955Z","iopub.status.idle":"2023-02-24T08:01:58.081722Z","shell.execute_reply.started":"2023-02-24T08:01:58.066913Z","shell.execute_reply":"2023-02-24T08:01:58.080621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make Pytorch Style Dataset Class","metadata":{}},{"cell_type":"code","source":"class MyDataset:\n    def __init__(self, df, use_da=0, data_type=\"train\"):\n        self.df = df.fillna(\"0\")\n        self.use_da = use_da\n        self.data_type = data_type\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n\n        r = self.df.iloc[i]\n        image_id = r.image_id\n        patient_id = r.patient_id\n        target = r.cancer\n        late = 1 if r.laterality==\"L\" else 0\n        view = 1 if r[\"view\"]==\"CC\" else 0\n        agen = int(r.age)//10\n        impl = int(r.implant)\n        img = cv2.imread(\"/kaggle/input/rsna-mammography-images-as-pngs/images_as_pngs_512/train_images_processed_512/%s/%s.png\"%(patient_id,image_id))\n        if img.shape[-1] == 1:\n            img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)\n        elif img.shape[-1] == 4:\n            img = img[...,:3]\n\n        img = img[...,::-1].astype(float) / 255.5\n\n        if self.use_da == 1:\n            img = horizontal_shift(img)\n            img = random_rotation(img)\n            img = random_flip(img)\n        if self.use_da == 2:\n            other = np.random.randint(2)\n            dfi = self.df[(self.df.cancer==other) & (self.df.view==r[\"view\"])]\n            dfir = dfi.iloc[np.random.randint(len(dfi))]\n            image_id = dfir.image_id\n            patient_id = dfir.patient_id\n            img2 = cv2.imread(\"/kaggle/input/rsna-mammography-images-as-pngs/images_as_pngs_512/train_images_processed_512/%s/%s.png\"%(patient_id,image_id))\n            if img2.shape[-1] == 1:\n                img2 = cv2.cvtColor(img2, cv2.COLOR_GRAY2BGR)\n            elif img2.shape[-1] == 4:\n                img2 = img2[...,:3]\n            img = random_paste(img, img2)\n            img = horizontal_shift(img)\n            img = random_rotation(img)\n            img = random_flip(img)\n            if target == 0:\n                target += other/2\n        if self.use_da == 3:\n            other = np.random.randint(2)\n            dfi = self.df[(self.df.cancer==other) & (self.df.view==r[\"view\"])]\n            dfir = dfi.iloc[np.random.randint(len(dfi))]\n            image_id = dfir.image_id\n            patient_id = dfir.patient_id\n            img2 = cv2.imread(\"/kaggle/input/rsna-mammography-images-as-pngs/images_as_pngs_512/train_images_processed_512/%s/%s.png\"%(patient_id,image_id))\n            if img2.shape[-1] == 1:\n                img2 = cv2.cvtColor(img2, cv2.COLOR_GRAY2BGR)\n            elif img2.shape[-1] == 4:\n                img2 = img2[...,:3]\n            img = horizontal_shift(img)\n            img = random_rotation(img)\n            img = random_flip(img)\n            img2 = horizontal_shift(img2)\n            img2 = random_rotation(img2)\n            img2 = random_flip(img2)\n            img = (img + img2) / 2 # MIX IN\n            if target!=other:\n                target = 1\n        \n        img = img.transpose((2,0,1))\n        img = torch.tensor(img)\n        ext = np.array([late,view,agen,impl], dtype=float) # np.float32 is not work in TPU\n        ext = torch.tensor(ext)\n        target = np.array([target], dtype=float) # np.float32 is not work in TPU\n        target = torch.tensor(target)\n        return img, ext, target","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.084959Z","iopub.execute_input":"2023-02-24T08:01:58.085286Z","iopub.status.idle":"2023-02-24T08:01:58.109526Z","shell.execute_reply.started":"2023-02-24T08:01:58.085251Z","shell.execute_reply":"2023-02-24T08:01:58.108499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds1 = MyDataset(df_train, use_da=1, data_type=\"train\")\ntrain_ds2 = MyDataset(df_train, use_da=2, data_type=\"train\")\ntrain_ds3 = MyDataset(df_train, use_da=3, data_type=\"train\")\ntest_ds = MyDataset(df_test, use_da=0, data_type=\"train\")\nvalid_ds = MyDataset(df_valid, use_da=0, data_type=\"train\")","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.111697Z","iopub.execute_input":"2023-02-24T08:01:58.11227Z","iopub.status.idle":"2023-02-24T08:01:58.133047Z","shell.execute_reply.started":"2023-02-24T08:01:58.11223Z","shell.execute_reply":"2023-02-24T08:01:58.132062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del df_train, df_test, df_valid\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.134384Z","iopub.execute_input":"2023-02-24T08:01:58.134702Z","iopub.status.idle":"2023-02-24T08:01:58.27803Z","shell.execute_reply.started":"2023-02-24T08:01:58.134652Z","shell.execute_reply":"2023-02-24T08:01:58.276775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Function","metadata":{}},{"cell_type":"code","source":"# 学習ループ\ndef train_one_epoch(epoch_no, data_loader, model, optimizer, device):\n    loss = get_LOSS()\n    model.train() # モデルを学習用に設定する\n    for X, E, y in tqdm(data_loader): # 画像を読み込んでtensorにする\n        X = X.to(device) # TPUを使うときはTPUメモリ上に乗せる\n        #E = E.to(device) # TPUを使うときはTPUメモリ上に乗せる\n        y = y.to(device) # TPUを使うときはTPUメモリ上に乗せる\n\n        # ニューラルネットワークを実行して損失値を求める\n        losses = loss(model(X), y)\n\n        # 新しいバッチ分の学習を行う\n        optimizer.zero_grad() # 一つ前の勾配をクリア\n        losses.backward() # 損失値を逆伝播させる\n        if GRAD_CLIP:\n            nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        xm.optimizer_step(optimizer) # 新しい勾配からパラメーターを更新する\n        \n        del X, y, losses\n        gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.279551Z","iopub.execute_input":"2023-02-24T08:01:58.280031Z","iopub.status.idle":"2023-02-24T08:01:58.290591Z","shell.execute_reply.started":"2023-02-24T08:01:58.279987Z","shell.execute_reply":"2023-02-24T08:01:58.288995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation Method\n\n**to use in TPU, caluclate by pytorch tensor**","metadata":{}},{"cell_type":"code","source":"def pfbeta_torch(labels, predictions, beta=1.0):\n    y_true_count = torch.sum(labels)\n    ctp = 0\n    cfp = 0\n\n    predictions = torch.clamp(predictions, min=0, max=1)\n    ctp = torch.sum(predictions * labels)\n    cfp = torch.sum(predictions * (1.0-labels))\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return 0","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.29265Z","iopub.execute_input":"2023-02-24T08:01:58.29351Z","iopub.status.idle":"2023-02-24T08:01:58.307249Z","shell.execute_reply.started":"2023-02-24T08:01:58.293434Z","shell.execute_reply":"2023-02-24T08:01:58.30571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation Function","metadata":{}},{"cell_type":"code","source":"# 評価ループ\ndef eval_one_epoch(epoch_no, data_loader, model, device, threashold, indexes):\n    #auroc = AUROC()\n    model.eval() # モデルを学習用に設定する\n    with torch.no_grad():\n        preds, trues = [], []\n        for X, E, y in tqdm(data_loader): # 画像を読み込んでtensorにする\n            X = X.to(device) # TPUを使うときはTPUメモリ上に乗せる\n            #E = E.to(device) # TPUを使うときはTPUメモリ上に乗せる\n            y = y.to(device) # TPUを使うときはTPUメモリ上に乗せる\n\n            # ニューラルネットワークを実行して損失値を求める\n            res = model(X)\n            preds.append(res[:,0])\n            trues.append(y[:,0].int())\n\n            del X, res\n            gc.collect()\n\n        # Get predicted scores for All Images\n        preds, trues = torch.cat(tuple(preds)), torch.cat(tuple(trues))\n        ccl, ccr, mlol, mlor = indexes\n\n        # Get predicted scores By View and laterality\n        preds_cc = torch.cat((preds[ccl], preds[ccr]))\n        preds_mlo = torch.cat((preds[mlol], preds[mlor]))\n        trues_cc = torch.cat((trues[ccl], trues[ccr]))\n        trues_mlo = torch.cat((trues[mlol], trues[mlor]))\n\n        # Prediction Score is patient has cancer, so it's maximum of CC and MLO view\n        preds = torch.max(preds_cc, preds_mlo)\n        trues = torch.max(trues_cc, trues_mlo)\n\n        # Make Threashold from test data\n        if threashold is None or threashold > 0:\n            best_score, best_threash = -1, 0\n            for i in range(1,1000,1):\n                t = i/1000\n                _preds = (preds > t).float() + 0.0001\n                _preds = torch.clamp(_preds, min=0, max=1)\n                score = pfbeta_torch(trues, _preds) #auroc(preds, trues) #-torch.mean((preds-trues)**2)\n                if score > best_score:\n                    best_score = score\n                    best_threash = t\n            score = best_score\n            threashold = best_threash\n        else: # caluclate score for validation data\n            t = threashold\n            preds = (preds > t).float() + 0.0001\n            preds = torch.clamp(preds, min=0, max=1)\n            score = pfbeta_torch(trues, preds) #auroc(preds, trues) #-torch.mean((preds-trues)**2)\n        del preds, trues\n        gc.collect()\n        return score, threashold","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.309368Z","iopub.execute_input":"2023-02-24T08:01:58.309854Z","iopub.status.idle":"2023-02-24T08:01:58.329266Z","shell.execute_reply.started":"2023-02-24T08:01:58.309799Z","shell.execute_reply":"2023-02-24T08:01:58.327994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training in TPU","metadata":{}},{"cell_type":"code","source":"# Updated\ndef _mp_fn(rank, flags):\n    # Acquires the (unique) TPU core corresponding to this process's index\n    device = xm.xla_device()\n\n    # Creates the (distributed) train sampler, which let this process only access\n    # its portion of the training dataset.\n    model = WDNet()\n    model.to(device)\n\n    optimizer = get_OPTIMIZER(model.parameters(), \n                              lr = flags['LR'])\n\n    scheduler = StepLR(optimizer,step_size = flags['EPOCHS']//3, gamma = 0.1, \n                                            last_epoch=-1)\n\n    xm.master_print('Training now...')\n    max_score = [-1,-1,-1]\n    for epoch in range(flags['EPOCHS']):\n\n        if epoch in flags['UPDATE_DA']:\n            nidx = flags['UPDATE_DA'].index(epoch)\n            train_sampler = DistributedSampler(dataset = flags['TRAIN_DS'][nidx],\n                                              num_replicas = xm.xrt_world_size(),\n                                              rank = xm.get_ordinal(),\n                                              shuffle = True)\n            train_dl = DataLoader(dataset = flags['TRAIN_DS'][nidx],\n                                  batch_size = flags['BATCH_SIZE'],\n                                  sampler = train_sampler,\n                                  num_workers = 0)\n            del train_sampler\n            gc.collect()\n\n\n        # Here comes our data loader for 8 cores.\n        # It takes famous 'DataLoader()' object and list of \n        # devices where data has to be sent.\n        # Calling 'per_device_loader()' on it will\n        # return the data loader for the particular device.\n        train_para_loader = pl.ParallelLoader(train_dl, \n                                              [device]).per_device_loader(device)\n\n        train_one_epoch(epoch, \n                        train_para_loader,\n                        model, \n                        optimizer, \n                        device)\n        scheduler.step()\n\n        del train_para_loader\n        gc.collect()\n        \n        if True:\n            test_sampler = DistributedSampler(dataset = flags['TEST_DS'],\n                                              num_replicas = xm.xrt_world_size(),\n                                              rank = xm.get_ordinal(),\n                                              shuffle = False)\n            test_dl = DataLoader(dataset = flags['TEST_DS'],\n                                  batch_size = flags['BATCH_SIZE_VAL'],\n                                  sampler = test_sampler,\n                                  num_workers = 0)\n            test_para_loader = pl.ParallelLoader(test_dl, \n                                                  [device]).per_device_loader(device)\n\n            del test_sampler, test_dl\n            gc.collect()\n            score1, threashold = eval_one_epoch(epoch, \n                            test_para_loader,\n                            model, \n                            device,\n                            threashold=None,\n                            indexes=flags['VALIDATE_INDEXES'][0])\n            del test_para_loader\n            gc.collect()\n\n            valid_sampler = DistributedSampler(dataset = flags['VALID_DS'],\n                                              num_replicas = xm.xrt_world_size(),\n                                              rank = xm.get_ordinal(),\n                                              shuffle = False)\n            valid_dl = DataLoader(dataset = flags['VALID_DS'],\n                                  batch_size = flags['BATCH_SIZE_VAL'],\n                                  sampler = valid_sampler,\n                                  num_workers = 0)\n            valid_para_loader = pl.ParallelLoader(valid_dl, \n                                                  [device]).per_device_loader(device)\n\n            del valid_sampler, valid_dl\n            gc.collect()\n            score2, _ = eval_one_epoch(epoch, \n                            valid_para_loader,\n                            model, \n                            device,\n                            threashold=threashold,\n                            indexes=flags['VALIDATE_INDEXES'][1])\n\n            del valid_para_loader\n            gc.collect()\n\n            xm.master_print(f\"epoch: {epoch} test score:{score1} valid score:{score2} threashold:{threashold}\")\n        \n        score = score2\n        for scoreidx in range(len(max_score)):\n            if score > max_score[scoreidx] and np.argmin(max_score) == scoreidx:\n                max_score[scoreidx] = score\n                #Saving the model, so that we can import it in the inference kernel.\n                xm.master_print(f\"save {epoch}epoch model to effnet-{scoreidx}.pth\")\n                xm.save(model.state_dict(), f\"myvit-{scoreidx}.pth\")\n                break\n        \n        gc.collect()\n        #xm.save(model.state_dict(), f\"effnet-{epoch}.pth\")\n\n    del model, optimizer, scheduler, max_score, train_dl\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.331124Z","iopub.execute_input":"2023-02-24T08:01:58.331427Z","iopub.status.idle":"2023-02-24T08:01:58.352028Z","shell.execute_reply.started":"2023-02-24T08:01:58.331391Z","shell.execute_reply":"2023-02-24T08:01:58.35086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"VALIDATE_INDEXES = [[test_index_CCL,\n                     test_index_CCR,\n                     test_index_MLOL,\n                     test_index_MLOR],[\n                     valid_index_CCL,\n                     valid_index_CCR,\n                     valid_index_MLOL,\n                     valid_index_MLOR]]","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.35342Z","iopub.execute_input":"2023-02-24T08:01:58.353903Z","iopub.status.idle":"2023-02-24T08:01:58.368478Z","shell.execute_reply.started":"2023-02-24T08:01:58.353867Z","shell.execute_reply":"2023-02-24T08:01:58.367018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FLAGS = {'TRAIN_DS': [train_ds1,train_ds2,train_ds3],\n         'TEST_DS': test_ds,\n         'VALID_DS': valid_ds,\n         'BATCH_SIZE': BATCH_SIZE,\n         'BATCH_SIZE_VAL': BATCH_SIZE_VAL,\n         'LR': 0.1,\n         'VALIDATE_INDEXES':VALIDATE_INDEXES,\n         'EPOCHS': NUM_EPOCHS,\n         'UPDATE_DA':UPDATE_DA_EPOCH}\n\ntry:\n    xmp.spawn(fn = _mp_fn, \n              args = (FLAGS,), \n              nprocs = xm.xrt_world_size())\nexcept:\n    pass","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:01:58.370169Z","iopub.execute_input":"2023-02-24T08:01:58.370572Z","iopub.status.idle":"2023-02-24T08:02:27.24276Z","shell.execute_reply.started":"2023-02-24T08:01:58.37052Z","shell.execute_reply":"2023-02-24T08:02:27.241745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%tb","metadata":{"execution":{"iopub.status.busy":"2023-02-24T08:02:27.244299Z","iopub.execute_input":"2023-02-24T08:02:27.245231Z","iopub.status.idle":"2023-02-24T08:02:27.250721Z","shell.execute_reply.started":"2023-02-24T08:02:27.245189Z","shell.execute_reply":"2023-02-24T08:02:27.249599Z"},"trusted":true},"execution_count":null,"outputs":[]}]}