{"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":"## Resources\nhttps://www.kaggle.com/code/younesselbrag/eda-training-using-pytorch-lighting-per-proces\nhttps://www.kaggle.com/code/tanreinama/training-efficientnet-with-tpu-in-rsna-screening\nhttps://www.kaggle.com/code/ooogunbiyi/image-embedding-eda-efficientnet","metadata":{}},{"cell_type":"code","source":"## Resources\n# https://www.kaggle.com/code/younesselbrag/eda-training-using-pytorch-lighting-per-proces\n# https://www.kaggle.com/code/tanreinama/training-efficientnet-with-tpu-in-rsna-screening\n# https://www.kaggle.com/code/ooogunbiyi/image-embedding-eda-efficientnet","metadata":{"execution":{"iopub.status.busy":"2022-12-04T05:32:48.171463Z","iopub.execute_input":"2022-12-04T05:32:48.171921Z","iopub.status.idle":"2022-12-04T05:32:48.194203Z","shell.execute_reply.started":"2022-12-04T05:32:48.171833Z","shell.execute_reply":"2022-12-04T05:32:48.193158Z"},"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 CosineAnnealingWarmRestarts\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":"2022-12-04T04:52:34.962115Z","iopub.execute_input":"2022-12-04T04:52:34.962366Z","iopub.status.idle":"2022-12-04T04:52:37.330235Z","shell.execute_reply.started":"2022-12-04T04:52:34.962325Z","shell.execute_reply":"2022-12-04T04:52:37.32895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 32\nBATCH_SIZE_VAL = 2\nNUM_EPOCHS = 18\nIMAGE_SIZE = 512\nNUM_USE = 2200\nTEST_RUN = True","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:37.332044Z","iopub.execute_input":"2022-12-04T04:52:37.332605Z","iopub.status.idle":"2022-12-04T04:52:37.337663Z","shell.execute_reply.started":"2022-12-04T04:52:37.332579Z","shell.execute_reply":"2022-12-04T04:52:37.336296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train metadata\ndf = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ndf[\"type\"] = \"train\"\nif TEST_RUN:\n    from sklearn.model_selection import train_test_split\n    NUM_USE = 220\n    NUM_EPOCHS = 3\n    df = pd.concat([df[df.cancer==1][:NUM_USE//2], df[df.cancer==0][:NUM_USE//2]], axis=0)\n    df_train, df_test = train_test_split(df, test_size=0.4, random_state=0)\n    df_test, df_valid = train_test_split(df_test, test_size=0.5, random_state=0)\n    del df\n    gc.collect()\nelse:\n    from sklearn.model_selection import train_test_split\n    df = pd.read_csv('../input/rsna-breast-cancer-detection/train.csv')\n    df[\"type\"] = \"train\"\n    dt = df[50000:]\n    df_testtrue = dt[dt.cancer==1]\n    df_testtrue, df_validtrue = train_test_split(df_testtrue, test_size=0.33, shuffle=True, random_state=0)\n\n    df_test = pd.concat([df_testtrue, dt[dt.cancer==0][:176]], axis=0)\n    df_valid = pd.concat([df_validtrue, dt[dt.cancer==0][176:1000]], axis=0)\n\n    dt = df[:49997]\n    df_train = pd.concat([dt[dt.cancer==1], dt[dt.cancer==0][:NUM_USE//2]], axis=0)\n    del dt, df\n    gc.collect()\nlen(df_train), len(df_test), len(df_valid)","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:37.339372Z","iopub.execute_input":"2022-12-04T04:52:37.339574Z","iopub.status.idle":"2022-12-04T04:52:37.618016Z","shell.execute_reply.started":"2022-12-04T04:52:37.339552Z","shell.execute_reply":"2022-12-04T04:52:37.616221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!mkdir ../temp","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:37.619631Z","iopub.execute_input":"2022-12-04T04:52:37.619929Z","iopub.status.idle":"2022-12-04T04:52:37.624686Z","shell.execute_reply.started":"2022-12-04T04:52:37.619887Z","shell.execute_reply":"2022-12-04T04:52:37.623255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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)","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:37.626205Z","iopub.execute_input":"2022-12-04T04:52:37.626582Z","iopub.status.idle":"2022-12-04T04:52:37.637385Z","shell.execute_reply.started":"2022-12-04T04:52:37.626549Z","shell.execute_reply":"2022-12-04T04:52:37.636618Z"},"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":"2022-12-04T04:52:37.638705Z","iopub.execute_input":"2022-12-04T04:52:37.639004Z","iopub.status.idle":"2022-12-04T04:52:37.659543Z","shell.execute_reply.started":"2022-12-04T04:52:37.638972Z","shell.execute_reply":"2022-12-04T04:52:37.658285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class WDNet(nn.Module):\n    def __init__(self):\n        super(WDNet, self).__init__()\n        conv = enet.EfficientNet.from_name('efficientnet-b6')\n        conv.load_state_dict(torch.load('../input/efficientnet-pytorch/efficientnet-b6-c76e70fd.pth'))\n        #conv._conv_stem.stride_ = 1\n        self.conv = conv\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc1 = nn.Linear(2308, 16)\n        self.fc2 = nn.Linear(16, 8)\n        self.fc3 = nn.Linear(8, 1)\n\n    def forward(self, input, ext):\n        x = self.conv.extract_features(input)\n        x = self.pool(x)\n        x = torch.flatten(x, start_dim=1)\n        x = torch.cat([x, ext], dim=-1)\n        x = torch.tanh(self.fc1(x))\n        x = torch.tanh(self.fc2(x))\n        x = torch.sigmoid(self.fc3(x))\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:37.660788Z","iopub.execute_input":"2022-12-04T04:52:37.661013Z","iopub.status.idle":"2022-12-04T04:52:37.673304Z","shell.execute_reply.started":"2022-12-04T04:52:37.66099Z","shell.execute_reply":"2022-12-04T04:52:37.67161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2022-12-04T04:52:37.674658Z","iopub.execute_input":"2022-12-04T04:52:37.674837Z","iopub.status.idle":"2022-12-04T04:52:37.863215Z","shell.execute_reply.started":"2022-12-04T04:52:37.674813Z","shell.execute_reply":"2022-12-04T04:52:37.862363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xm.get_xla_supported_devices()","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:37.865728Z","iopub.execute_input":"2022-12-04T04:52:37.865944Z","iopub.status.idle":"2022-12-04T04:52:49.338053Z","shell.execute_reply.started":"2022-12-04T04:52:37.865918Z","shell.execute_reply":"2022-12-04T04:52:49.337222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xm.xrt_world_size()","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:49.339312Z","iopub.execute_input":"2022-12-04T04:52:49.339594Z","iopub.status.idle":"2022-12-04T04:52:49.355131Z","shell.execute_reply.started":"2022-12-04T04:52:49.339567Z","shell.execute_reply":"2022-12-04T04:52:49.353602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Data Augumentation\n\ndef cut_patch(img):\n    img_height, img_width, _ = img.shape\n    top = random.randrange(0, round(img_height))\n    bottom = top + random.randrange(round(img_height*0.05),\n                                            round(img_height*0.15))\n    left = random.randrange(0, round(img_width))\n    right = left + random.randrange(round(img_width*0.05),\n                                            round(img_width*0.15))\n    if (bottom - top) % 2 == 1:\n        bottom -= 1\n    if (right - left) % 2 == 1:\n        right -= 1\n    return img[top:bottom, left:right, :]\n\ndef paste_patch(img, patch, rot, ratio):\n    img_height, img_width, _ = img.shape\n    patch_height, patch_width, _ = patch.shape\n    width_half = round(img_width / 2)\n    height_half = round(img_height / 2)\n    trans_x = random.randrange(-width_half, width_half)\n    trans_y = random.randrange(-height_half, height_half)\n    patch_h_center = round(patch_height / 2)\n    patch_w_center = round(patch_width / 2)\n    img_h_center = round(img_height / 2)\n    img_w_center = round(img_width / 2)\n    top = round((img_height - patch_height) / 2)\n    bottom = round((img_height + patch_height) / 2)\n    left = round((img_width - patch_width) / 2)\n    right = round((img_width + patch_width) / 2)\n    # paste on center\n    tmp_img = np.zeros((img_height, img_width, 3), np.uint8)\n    tmp_img[top:bottom, left:right, :] = patch\n    # rotation and expansion\n    M = cv2.getRotationMatrix2D((img_w_center, img_h_center), rot, ratio)\n    tmp_img = cv2.warpAffine(tmp_img, M, (img_width, img_height))\n    # translation\n    M = np.float32([[1, 0, trans_x], [0, 1, trans_y]])\n    tmp_img = cv2.warpAffine(tmp_img, M, (img_width, img_height))\n    # make mask of patch\n    imggray = cv2.cvtColor(tmp_img, cv2.COLOR_BGR2GRAY)\n    ret, mask = cv2.threshold(imggray, 10, 255, cv2.THRESH_BINARY)\n    mask_inv = cv2.bitwise_not(mask)\n    # cut the mask from original image\n    back = cv2.bitwise_and(img, img, mask=mask_inv)\n    cut = cv2.bitwise_and(tmp_img, tmp_img, mask = mask)\n    # paste(combine original and patch)\n    paste = cv2.add(back, cut)\n    return paste\n\ndef random_paste(img, img2):\n    patch = cut_patch(img2)\n    angle = int(random.uniform(-30, 30))\n    ratio = np.random.random() / 2\n    return paste_patch(img, patch, angle, ratio)\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":"2022-12-04T04:52:49.356886Z","iopub.execute_input":"2022-12-04T04:52:49.357078Z","iopub.status.idle":"2022-12-04T04:52:49.375483Z","shell.execute_reply.started":"2022-12-04T04:52:49.357056Z","shell.execute_reply":"2022-12-04T04:52:49.374452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        if self.use_da >= 1:\n            img = random_flip(img)\n        elif self.use_da >= 2:\n            img = horizontal_shift(img)\n            img = random_rotation(img)\n            img = random_flip(img)\n        elif self.use_da >= 3:\n            dfi = self.df[self.df.target==target]\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(img2)\n            img = horizontal_shift(img)\n            img = random_rotation(img)\n            img = random_flip(img)\n        \n        img = img[...,::-1]\n        img = img.transpose((2,0,1)).astype(float) / 255.5\n        img = torch.tensor(img)\n        ext = np.array([late,view,agen,impl], dtype=float)\n        ext = torch.tensor(ext)\n        target = np.array([target], dtype=float)\n        target = torch.tensor(target)\n        return img, ext, target","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:49.376531Z","iopub.execute_input":"2022-12-04T04:52:49.376705Z","iopub.status.idle":"2022-12-04T04:52:49.393224Z","shell.execute_reply.started":"2022-12-04T04:52:49.376683Z","shell.execute_reply":"2022-12-04T04:52:49.392561Z"},"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":"2022-12-04T04:52:49.394083Z","iopub.execute_input":"2022-12-04T04:52:49.394246Z","iopub.status.idle":"2022-12-04T04:52:49.413325Z","shell.execute_reply.started":"2022-12-04T04:52:49.394226Z","shell.execute_reply":"2022-12-04T04:52:49.412644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del df_train, df_test\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:49.41449Z","iopub.execute_input":"2022-12-04T04:52:49.414723Z","iopub.status.idle":"2022-12-04T04:52:49.526009Z","shell.execute_reply.started":"2022-12-04T04:52:49.414694Z","shell.execute_reply":"2022-12-04T04:52:49.524845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 学習ループ\ndef train_one_epoch(epoch_no, data_loader, model, optimizer, device):\n    loss = FocalLoss()\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, E), y)\n\n        # 新しいバッチ分の学習を行う\n        optimizer.zero_grad() # 一つ前の勾配をクリア\n        losses.backward() # 損失値を逆伝播させる\n        xm.optimizer_step(optimizer) # 新しい勾配からパラメーターを更新する\n        \n        del X, y, losses\n        gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:52:49.527934Z","iopub.execute_input":"2022-12-04T04:52:49.528615Z","iopub.status.idle":"2022-12-04T04:52:49.53662Z","shell.execute_reply.started":"2022-12-04T04:52:49.528588Z","shell.execute_reply":"2022-12-04T04:52:49.535765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2022-12-04T05:00:19.873902Z","iopub.execute_input":"2022-12-04T05:00:19.874175Z","iopub.status.idle":"2022-12-04T05:00:19.882027Z","shell.execute_reply.started":"2022-12-04T05:00:19.87415Z","shell.execute_reply":"2022-12-04T05:00:19.880419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 評価ループ\ndef eval_one_epoch(epoch_no, data_loader, model, device, threashold):\n    #auroc = AUROC()\n    model.eval() # モデルを学習用に設定する\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, E)\n        preds.append(res)\n        trues.append(y.int())\n        \n        del X, res\n        gc.collect()\n    preds, trues = torch.cat(tuple(preds)), torch.cat(tuple(trues))\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:\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":"2022-12-04T05:00:21.742988Z","iopub.execute_input":"2022-12-04T05:00:21.743221Z","iopub.status.idle":"2022-12-04T05:00:21.753999Z","shell.execute_reply.started":"2022-12-04T05:00:21.743196Z","shell.execute_reply":"2022-12-04T05:00:21.752933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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 = optim.Adam(model.parameters(), \n                              lr = flags['LR'])\n\n    scheduler = CosineAnnealingWarmRestarts(optimizer, \n                                            T_0 = flags['EPOCHS'], \n                                            T_mult=1, \n                                            eta_min=0.0001, \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 % (flags['EPOCHS']//3) == 0:\n            train_sampler = DistributedSampler(dataset = flags['TRAIN_DS'][epoch // (flags['EPOCHS']//3)],\n                                              num_replicas = xm.xrt_world_size(),\n                                              rank = xm.get_ordinal(),\n                                              shuffle = True)\n            train_dl = DataLoader(dataset = flags['TRAIN_DS'][epoch // (flags['EPOCHS']//3)],\n                                  batch_size = flags['BATCH_SIZE'],\n                                  sampler = train_sampler,\n                                  num_workers = 0)\n\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        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        score, threashold = eval_one_epoch(epoch, \n                        test_para_loader,\n                        model, \n                        device,\n                        threashold=None)\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        score, _ = eval_one_epoch(epoch, \n                        valid_para_loader,\n                        model, \n                        device,\n                        threashold=threashold)\n        \n        del valid_para_loader\n        gc.collect()\n                \n        xm.master_print(f\"epoch: {epoch} test score:{score} valid score:{score} threashold:{threashold}\")\n        #\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\"effnet-{scoreidx}.pth\")\n        #        break\n        #\n        #del test_dl, test_para_loader, score\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":"2022-12-04T05:00:36.171776Z","iopub.execute_input":"2022-12-04T05:00:36.172022Z","iopub.status.idle":"2022-12-04T05:00:36.189132Z","shell.execute_reply.started":"2022-12-04T05:00:36.172Z","shell.execute_reply":"2022-12-04T05:00:36.187954Z"},"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.001,\n         'EPOCHS': NUM_EPOCHS}\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":"2022-12-04T05:00:40.28033Z","iopub.execute_input":"2022-12-04T05:00:40.280607Z","iopub.status.idle":"2022-12-04T05:03:49.295752Z","shell.execute_reply.started":"2022-12-04T05:00:40.280577Z","shell.execute_reply":"2022-12-04T05:03:49.294333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%tb","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:59:29.457529Z","iopub.execute_input":"2022-12-04T04:59:29.457919Z","iopub.status.idle":"2022-12-04T04:59:29.463008Z","shell.execute_reply.started":"2022-12-04T04:59:29.457895Z","shell.execute_reply":"2022-12-04T04:59:29.462054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2022-12-04T04:59:29.464285Z","iopub.execute_input":"2022-12-04T04:59:29.466691Z","iopub.status.idle":"2022-12-04T04:59:29.780681Z","shell.execute_reply.started":"2022-12-04T04:59:29.466627Z","shell.execute_reply":"2022-12-04T04:59:29.77942Z"},"trusted":true},"execution_count":null,"outputs":[]}]}