{"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":"# import","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/pytorch112-cu113/{torch-1.12.1+cu113-cp37-cp37m-linux_x86_64.whl,torchvision-0.13.1+cu113-cp37-cp37m-linux_x86_64.whl}","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:13:13.662293Z","iopub.execute_input":"2023-02-27T10:13:13.662632Z","iopub.status.idle":"2023-02-27T10:14:49.394379Z","shell.execute_reply.started":"2023-02-27T10:13:13.662596Z","shell.execute_reply":"2023-02-27T10:14:49.393167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp /kaggle/input/nvjpeg2k/nvjpeg2k.so ./","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:14:49.397138Z","iopub.execute_input":"2023-02-27T10:14:49.397518Z","iopub.status.idle":"2023-02-27T10:14:50.401797Z","shell.execute_reply.started":"2023-02-27T10:14:49.397475Z","shell.execute_reply":"2023-02-27T10:14:50.400454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/dicomsdl-offline-installer/dicomsdl-0.109.1-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:14:50.404232Z","iopub.execute_input":"2023-02-27T10:14:50.404644Z","iopub.status.idle":"2023-02-27T10:15:02.102133Z","shell.execute_reply.started":"2023-02-27T10:14:50.404603Z","shell.execute_reply":"2023-02-27T10:15:02.100869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    import pylibjpeg\nexcept:\n    !pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:02.105812Z","iopub.execute_input":"2023-02-27T10:15:02.106238Z","iopub.status.idle":"2023-02-27T10:15:14.603356Z","shell.execute_reply.started":"2023-02-27T10:15:02.106191Z","shell.execute_reply":"2023-02-27T10:15:14.602108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip uninstall -y timm","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:14.605061Z","iopub.execute_input":"2023-02-27T10:15:14.605715Z","iopub.status.idle":"2023-02-27T10:15:16.685167Z","shell.execute_reply.started":"2023-02-27T10:15:14.605675Z","shell.execute_reply":"2023-02-27T10:15:16.683898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\n\nsys.path.append(\"../input/timmmaster/\")\n\nimport argparse\nimport gc\nimport glob\n\nimport albumentations as A\nimport cv2\nimport dicomsdl\nimport joblib\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport pytorch_lightning as pl\nimport re\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport torchvision\nimport wandb\nimport yaml\nfrom concurrent.futures import ProcessPoolExecutor\nfrom joblib import Parallel, delayed\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nfrom pytorch_lightning import Trainer, seed_everything\nfrom pytorch_lightning.loggers import WandbLogger\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom tqdm import tqdm\n\nfrom pydicom.filebase import DicomBytesIO\nimport pydicom","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:16.688127Z","iopub.execute_input":"2023-02-27T10:15:16.688767Z","iopub.status.idle":"2023-02-27T10:15:29.488984Z","shell.execute_reply.started":"2023-02-27T10:15:16.688722Z","shell.execute_reply":"2023-02-27T10:15:29.487796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# cfg","metadata":{}},{"cell_type":"code","source":"with open(\"/kaggle/input/rsna-exp296/exp296.yaml\", encoding=\"utf-8\") as f:\n    cfg1 = yaml.safe_load(f)\ncfg1[\"general\"][\"wandb_desabled\"] = True\ncfg1[\"general\"][\"input_path\"] = \"/kaggle/input/rsna-breast-cancer-detection\"\ncfg1[\"general\"][\"output_path\"] = \"/kaggle/input\"\ncfg1[\"general\"][\"save_name\"] = \"rsna-exp296\"\ncfg1[\"model\"][\"model_save\"] = False\ncfg1[\"pl_params\"][\"enable_checkpointing\"] = False\ncfg1[\"pl_params\"][\"precision\"] = 16\ncfg1[\"general\"][\"fold\"] = None\ncfg1[\"tta\"] = False\ncfg1[\"general\"][\"cv\"] = False\n\nwith open(\"/kaggle/input/rsna-exp302/exp302.yaml\", encoding=\"utf-8\") as f:\n    cfg2 = yaml.safe_load(f)\ncfg2[\"general\"][\"wandb_desabled\"] = True\ncfg2[\"general\"][\"input_path\"] = \"/kaggle/input/rsna-breast-cancer-detection\"\ncfg2[\"general\"][\"output_path\"] = \"/kaggle/input\"\ncfg2[\"general\"][\"save_name\"] = \"rsna-exp302\"\ncfg2[\"model\"][\"model_save\"] = False\ncfg2[\"pl_params\"][\"enable_checkpointing\"] = False\ncfg2[\"pl_params\"][\"precision\"] = 16\ncfg2[\"general\"][\"fold\"] = None\ncfg2[\"tta\"] = False\ncfg2[\"general\"][\"cv\"] = False\n\nwith open(\"/kaggle/input/rsna-exp288/exp288.yaml\", encoding=\"utf-8\") as f:\n    cfg3 = yaml.safe_load(f)\ncfg3[\"general\"][\"wandb_desabled\"] = True\ncfg3[\"general\"][\"input_path\"] = \"/kaggle/input/rsna-breast-cancer-detection\"\ncfg3[\"general\"][\"output_path\"] = \"/kaggle/input\"\ncfg3[\"general\"][\"save_name\"] = \"rsna-exp288\"\ncfg3[\"model\"][\"model_save\"] = False\ncfg3[\"pl_params\"][\"enable_checkpointing\"] = False\ncfg3[\"pl_params\"][\"precision\"] = 16\ncfg3[\"general\"][\"fold\"] = None\ncfg3[\"tta\"] = False\ncfg3[\"general\"][\"cv\"] = False\n\ncfg_list = [cfg1, cfg2, cfg3]","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:29.494344Z","iopub.execute_input":"2023-02-27T10:15:29.496767Z","iopub.status.idle":"2023-02-27T10:15:29.585311Z","shell.execute_reply.started":"2023-02-27T10:15:29.49672Z","shell.execute_reply":"2023-02-27T10:15:29.58425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"GPU = True\nUSE_DATA = \"test\"\nBINARIZED = True\nAGG = \"top3_mean\" # mean or max\nTHR = 0.44","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:29.590166Z","iopub.execute_input":"2023-02-27T10:15:29.592862Z","iopub.status.idle":"2023-02-27T10:15:29.599647Z","shell.execute_reply.started":"2023-02-27T10:15:29.592817Z","shell.execute_reply":"2023-02-27T10:15:29.598646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if GPU:\n    import nvjpeg2k\n    decoder = nvjpeg2k.Decoder()\nelse:\n    decoder = None","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:29.604685Z","iopub.execute_input":"2023-02-27T10:15:29.607558Z","iopub.status.idle":"2023-02-27T10:15:33.394308Z","shell.execute_reply.started":"2023-02-27T10:15:29.607508Z","shell.execute_reply":"2023-02-27T10:15:33.39326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Preparation","metadata":{}},{"cell_type":"code","source":"if USE_DATA == \"train\":\n    csv_file = \"/kaggle/input/rsna-breast-cancer-detection/train.csv\"\n    dcm_dir  = \"/kaggle/input/rsna-breast-cancer-detection/train_images\"\n\nif USE_DATA == \"test\":\n    csv_file = \"/kaggle/input/rsna-breast-cancer-detection/test.csv\"\n    dcm_dir  = \"/kaggle/input/rsna-breast-cancer-detection/test_images\"\n\nSIZE = 1536\nVOI = cfg1[\"task\"][\"voi\"] # True or False\nBIT = cfg1[\"task\"][\"bit\"] # 8, 16\nEXTENSION = cfg1[\"task\"][\"extend\"]\nN_JOBS = 2 if GPU else 4 # gpu 2, cpu 4\n\nif VOI:\n    SAVE_FOLDER = f\"/kaggle/tmp/output/{EXTENSION}_{SIZE}_{BIT}bit_voi/\"\nelse:\n    SAVE_FOLDER = f\"/kaggle/tmp/output/{EXTENSION}_{SIZE}_{BIT}bit/\"\n    \nos.makedirs(SAVE_FOLDER, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:33.399808Z","iopub.execute_input":"2023-02-27T10:15:33.400698Z","iopub.status.idle":"2023-02-27T10:15:33.408791Z","shell.execute_reply.started":"2023-02-27T10:15:33.400657Z","shell.execute_reply":"2023-02-27T10:15:33.407628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_transfer_syntax_uid(df, dcm_dir):\n    machine_id_to_transfer = {}\n    machine_id = df.machine_id.unique()\n    for i in machine_id:\n        d = df[df.machine_id == i].iloc[0]\n        f = f\"{dcm_dir}/{d.patient_id}/{d.image_id}.dcm\"\n        dicom = pydicom.dcmread(f)\n        machine_id_to_transfer[i] = dicom.file_meta.TransferSyntaxUID\n    return machine_id_to_transfer","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:33.410269Z","iopub.execute_input":"2023-02-27T10:15:33.411135Z","iopub.status.idle":"2023-02-27T10:15:33.425244Z","shell.execute_reply.started":"2023-02-27T10:15:33.411094Z","shell.execute_reply":"2023-02-27T10:15:33.424289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process(f, size=512, save_folder=\"\", extension=\"png\", force_slow=False):\n    patient = f.split('/')[-2]\n    image = f.split('/')[-1][:-4]\n    \n    meta = pydicom.dcmread(f, stop_before_pixels=True)\n    if meta.file_meta.TransferSyntaxUID == \"1.2.840.10008.1.2.4.90\":\n        if not force_slow:\n            #print(\"jpeg2k\")\n            with open(f, \"rb\") as _f:\n                raw = DicomBytesIO(_f.read())\n                ds = pydicom.dcmread(raw)\n            offset = ds.PixelData.find(b\"\\x00\\x00\\x00\\x0C\")\n            hackedbitstream = bytearray()\n            hackedbitstream.extend(ds.PixelData[offset:])\n            img = decoder.decode(hackedbitstream)\n        else:\n            #print(\"pydicom\")\n            img = pydicom.dcmread(f).pixel_array\n    else:\n        #print(\"dicomsdl\")\n        dicom = dicomsdl.open(f)\n        img = dicom.pixelData()\n    \n    if VOI:\n        img = apply_voi_lut(img, meta)\n\n    img = (img - img.min()) / (img.max() - img.min())\n\n    if meta.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n\n    #img = cv2.resize(img, (size, size))\n    \n    if BIT == 16:\n        img = (img * 65535).astype(np.uint16)\n    else:\n        img = (img * 255).astype(np.uint8)\n\n    fit_image(img, patient, image, extension)","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:33.426845Z","iopub.execute_input":"2023-02-27T10:15:33.427353Z","iopub.status.idle":"2023-02-27T10:15:33.440227Z","shell.execute_reply.started":"2023-02-27T10:15:33.427316Z","shell.execute_reply":"2023-02-27T10:15:33.43918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def image_resize(image, width = None, height = None, inter = cv2.INTER_LINEAR):\n\n    dim = None\n    (h, w) = image.shape[:2]\n\n    if width is None and height is None:\n        return image\n\n    if width is None:\n        r = height / float(h)\n        dim = (int(w * r), height)\n    else:\n        r = width / float(w)\n        dim = (width, int(h * r))\n    resized = cv2.resize(image, dim, interpolation = inter)\n\n    return resized","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:33.441788Z","iopub.execute_input":"2023-02-27T10:15:33.442187Z","iopub.status.idle":"2023-02-27T10:15:33.454282Z","shell.execute_reply.started":"2023-02-27T10:15:33.442146Z","shell.execute_reply":"2023-02-27T10:15:33.453126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fit_image(X, patient_id, im_id, extension, return_image=False):\n    #X = cv2.imread(fname)\n    \n    # Some images have narrow exterior \"frames\" that complicate selection of the main data. Cutting off the frame\n    X = X[5:-5, 5:-5]\n    \n    # regions of non-empty pixels\n    output= cv2.connectedComponentsWithStats((X > 20).astype(np.uint8), 8, cv2.CV_32S)\n\n    # stats.shape == (N, 5), where N is the number of regions, 5 dimensions correspond to:\n    # left, top, width, height, area_size\n    stats = output[2]\n\n    # finding max area which always corresponds to the breast data. \n    idx = stats[1:, 4].argmax() + 1\n    x1, y1, w, h = stats[idx][:4]\n    x2 = x1 + w\n    y2 = y1 + h\n    \n    # cutting out the breast data\n    X_fit = X[y1: y2, x1: x2]\n    \n    #patient_id, im_id = re.findall(\"(\\d+)_(\\d+).png\", os.path.basename(fname))[0]\n    image = X_fit\n    \n    shape = image.shape\n    if shape[0] > shape[1]:\n        image = image_resize(image, height=SIZE)\n    else:\n        image = image_resize(image, width=SIZE)\n        \n    if return_image:\n        return image\n    else:\n        cv2.imwrite(SAVE_FOLDER + f\"{patient_id}_{im_id}.{extension}\", image)\n\n#def fit_all_images(all_images):\n#    with ProcessPoolExecutor(4) as p:\n#        for i in tqdm(p.map(fit_image, all_images), total=len(all_images)):\n#            pass","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:33.455674Z","iopub.execute_input":"2023-02-27T10:15:33.457391Z","iopub.status.idle":"2023-02-27T10:15:33.467084Z","shell.execute_reply.started":"2023-02-27T10:15:33.457345Z","shell.execute_reply":"2023-02-27T10:15:33.466124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(csv_file)\nmachine_id_to_transfer = make_transfer_syntax_uid(test_df, dcm_dir)\ntest_df.loc[:, \"i\"] = np.arange(len(test_df))\ntest_df.loc[:, \"TransferSyntaxUID\"] = test_df.machine_id.map(machine_id_to_transfer)    \n\nprint(\"test_df\", test_df.shape)\nprint(test_df)\nprint(\"\")","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:33.47014Z","iopub.execute_input":"2023-02-27T10:15:33.470743Z","iopub.status.idle":"2023-02-27T10:15:33.624218Z","shell.execute_reply.started":"2023-02-27T10:15:33.470715Z","shell.execute_reply":"2023-02-27T10:15:33.623016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"j2k_df = test_df[test_df.TransferSyntaxUID == '1.2.840.10008.1.2.4.90'].reset_index(drop=True)\nnon_j2k_df = test_df[test_df.TransferSyntaxUID != '1.2.840.10008.1.2.4.90'].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:33.625976Z","iopub.execute_input":"2023-02-27T10:15:33.626355Z","iopub.status.idle":"2023-02-27T10:15:33.63546Z","shell.execute_reply.started":"2023-02-27T10:15:33.626316Z","shell.execute_reply":"2023-02-27T10:15:33.633818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if GPU:\n    print(\"jpeg2k\")\n    for t, d in tqdm(j2k_df.iterrows()):\n        dcm_file = f\"{dcm_dir}/{d.patient_id}/{d.image_id}.dcm\"\n        process(dcm_file, size=SIZE, save_folder=SAVE_FOLDER, extension=EXTENSION, force_slow=False)\nelse:\n    print(\"pydicom\")\n    Parallel(n_jobs=N_JOBS, backend=\"multiprocessing\")(\n        delayed(process)(f\"{dcm_dir}/{d.patient_id}/{d.image_id}.dcm\", size=SIZE, save_folder=SAVE_FOLDER, extension=EXTENSION, force_slow=True)\n        for t, d in tqdm(j2k_df.iterrows())\n    )","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:33.638238Z","iopub.execute_input":"2023-02-27T10:15:33.639444Z","iopub.status.idle":"2023-02-27T10:15:34.976396Z","shell.execute_reply.started":"2023-02-27T10:15:33.639402Z","shell.execute_reply":"2023-02-27T10:15:34.975272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Parallel(n_jobs=N_JOBS, backend=\"multiprocessing\")(\n    delayed(process)(f\"{dcm_dir}/{d.patient_id}/{d.image_id}.dcm\", size=SIZE, save_folder=SAVE_FOLDER, extension=EXTENSION)\n    for t, d in tqdm(non_j2k_df.iterrows())\n)","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:34.978525Z","iopub.execute_input":"2023-02-27T10:15:34.979284Z","iopub.status.idle":"2023-02-27T10:15:35.61154Z","shell.execute_reply.started":"2023-02-27T10:15:34.979241Z","shell.execute_reply":"2023-02-27T10:15:35.610042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.listdir(SAVE_FOLDER)","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.616698Z","iopub.execute_input":"2023-02-27T10:15:35.617677Z","iopub.status.idle":"2023-02-27T10:15:35.635498Z","shell.execute_reply.started":"2023-02-27T10:15:35.617606Z","shell.execute_reply":"2023-02-27T10:15:35.633893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# def class and function","metadata":{}},{"cell_type":"markdown","source":"### model","metadata":{}},{"cell_type":"code","source":"def pfbeta(labels, predictions, beta=1):\n    y_true_count = 0\n    ctp = 0\n    cfp = 0\n\n    for idx in range(len(labels)):\n        prediction = min(max(predictions[idx], 0), 1)\n        if labels[idx]:\n            y_true_count += 1\n            ctp += prediction\n        else:\n            cfp += prediction\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp + 1e-7)\n    c_recall = ctp / (y_true_count + 1e-7)\n    if c_precision > 0 and c_recall > 0:\n        result = (\n            (1 + beta_squared)\n            * (c_precision * c_recall)\n            / (beta_squared * c_precision + c_recall + 1e-7)\n        )\n        return result\n    else:\n        return 0\n\ndef optimal_f1(labels, predictions):\n    thres = np.arange(0, 1, 0.01)\n    f1s = [pfbeta(labels, predictions > thr) for thr in thres]\n    idx = np.argmax(f1s)\n    return f1s[idx], thres[idx]\n    \nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=1, gamma=2, logits=False, reduce=True, pos_weight=None):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.logits = logits\n        self.reduce = reduce\n        self.pos_weight = pos_weight\n\n    def forward(self, inputs, targets):\n        if self.logits:\n            BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduce=False, pos_weight=self.pos_weight)\n        else:\n            BCE_loss = F.binary_cross_entropy(inputs, targets, reduce=False)\n        pt = torch.exp(-BCE_loss)\n        F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss\n\n        if self.reduce:\n            return torch.mean(F_loss)\n        else:\n            return F_loss\n        \nclass ArcMarginProduct(nn.Module):\n    r\"\"\"Implement of large margin arc distance: :\n    Args:\n        in_features: size of each input sample\n        out_features: size of each output sample\n        s: norm of input feature\n        m: margin\n        cos(theta + m)\n    \"\"\"\n\n    def __init__(\n        self,\n        in_features: int,\n        out_features: int,\n        s: float,\n        m: float,\n        easy_margin: bool,\n        ls_eps: float,\n    ):\n        super(ArcMarginProduct, self).__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.s = s\n        self.m = m\n        self.ls_eps = ls_eps  # label smoothing\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.easy_margin = easy_margin\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n\n    def forward(self, input: torch.Tensor, label: torch.Tensor, device: str = \"cuda\") -> torch.Tensor:\n        # --------------------------- cos(theta) & phi(theta) ---------------------\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        # Enable 16 bit precision\n        cosine = cosine.to(torch.float32)\n\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n        phi = cosine * self.cos_m - sine * self.sin_m\n        if self.easy_margin:\n            phi = torch.where(cosine > 0, phi, cosine)\n        else:\n            phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        # --------------------------- convert label to one-hot ---------------------\n        # one_hot = torch.zeros(cosine.size(), requires_grad=True, device='cuda')\n        one_hot = torch.zeros(cosine.size(), device=device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        if self.ls_eps > 0:\n            one_hot = (1 - self.ls_eps) * one_hot + self.ls_eps / self.out_features\n        # -------------torch.where(out_i = {x_i if condition_i else y_i) ------------\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n\n        return output","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.637581Z","iopub.execute_input":"2023-02-27T10:15:35.638196Z","iopub.status.idle":"2023-02-27T10:15:35.671594Z","shell.execute_reply.started":"2023-02-27T10:15:35.638151Z","shell.execute_reply":"2023-02-27T10:15:35.670152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CNNModel(pl.LightningModule):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        if cfg[\"arcface\"] is None or cfg[\"task\"][\"aux_target\"] is None:\n            num_classes = 1\n        else:\n            num_classes = None\n        self.model = timm.create_model(\n            model_name=cfg[\"model\"][\"model_name\"],\n            pretrained=False, #cfg[\"model\"][\"pretrained\"],\n            in_chans=cfg[\"model\"][\"in_chans\"],\n            num_classes=num_classes,\n            drop_rate=cfg[\"model\"][\"drop_rate\"],\n            drop_path_rate=cfg[\"model\"][\"drop_path_rate\"],\n        )\n        \n        if cfg[\"model\"][\"criterion\"] == \"FocalLoss\":\n            self.criterion = FocalLoss(logits=True, pos_weight=cfg[\"model\"][\"loss_weights\"])\n        else:\n            self.criterion = nn.__dict__[cfg[\"model\"][\"criterion\"]](pos_weight=cfg[\"model\"][\"loss_weights\"])\n            \n        if cfg[\"arcface\"] is not None:\n            self.embedding = nn.Linear(self.model.get_classifier().in_features, cfg[\"arcface\"][\"params\"][\"in_features\"])\n            self.arc = ArcMarginProduct(**cfg[\"arcface\"][\"params\"])\n            self.criterion = nn.CrossEntropyLoss()\n            self.model.reset_classifier(num_classes=0, global_pool=\"avg\")\n        elif cfg[\"task\"][\"aux_target\"] is not None:\n            self.nn_cancer = torch.nn.Sequential(torch.nn.Linear(self.model.get_classifier().in_features, 1))\n            self.nn_aux = torch.nn.ModuleList([torch.nn.Linear(self.model.get_classifier().in_features, n) for n in cfg[\"task\"][\"aux_target_nclasses\"]])\n            self.model.reset_classifier(num_classes=0, global_pool=\"avg\")\n        \n        if cfg[\"model\"][\"grad_checkpointing\"]:\n            print(\"grad_checkpointing true\")\n            self.model.set_grad_checkpointing(enable=True)\n\n    def forward(self, X):\n        if self.cfg[\"arcface\"] is not None:\n            X = self.model(X)\n            outputs = self.embedding(X)\n        elif self.cfg[\"task\"][\"aux_target\"] is not None:\n            X = self.model(X)\n            cancer = self.nn_cancer(X) #.squeeze()\n            aux = []\n            for nn in self.nn_aux:\n                aux.append(nn(X)) #.squeeze()\n            return cancer, aux\n        else:\n            outputs = self.model(X)\n        return outputs\n    \n    def rand_bbox(self, size, lam):\n        W = size[2]\n        H = size[3]\n        cut_rat = np.sqrt(1. - lam)\n        cut_w = np.int(W * cut_rat)\n        cut_h = np.int(H * cut_rat)\n\n        # uniform\n        cx = np.random.randint(W)\n        cy = np.random.randint(H)\n\n        bbx1 = np.clip(cx - cut_w // 2, 0, W)\n        bby1 = np.clip(cy - cut_h // 2, 0, H)\n        bbx2 = np.clip(cx + cut_w // 2, 0, W)\n        bby2 = np.clip(cy + cut_h // 2, 0, H)\n        return bbx1, bby1, bbx2, bby2\n\n    def cutmix_data(self, x, y, alpha=1.0):\n        indices = torch.randperm(x.size(0))\n        shuffled_data = x[indices]\n        shuffled_target = y[indices]\n\n        lam = np.clip(np.random.beta(alpha, alpha),0.3,0.4)\n        bbx1, bby1, bbx2, bby2 = self.rand_bbox(x.size(), lam)\n        new_data = x.clone()\n        new_data[:, :, bby1:bby2, bbx1:bbx2] = x[indices, :, bby1:bby2, bbx1:bbx2]\n        # adjust lambda to exactly match pixel ratio\n        lam = 1 - ((bbx2 - bbx1) * (bby2 - bby1) / (x.size()[-1] * x.size()[-2]))\n\n        return new_data, y, shuffled_target, lam\n    \n    def mixup_data(self, x, y, alpha=1.0, return_index=False):\n        if alpha > 0:\n            lam = np.random.beta(alpha, alpha)\n        else:\n            lam = 1\n\n        batch_size = x.size()[0]\n        index = torch.randperm(batch_size)\n        mixed_x = lam * x + (1 - lam) * x[index, :]\n        y_a, y_b = y, y[index]\n        \n        if return_index:\n            return mixed_x, y_a, y_b, lam, index\n        else:\n            return mixed_x, y_a, y_b, lam\n    \n    def mix_criterion(self, pred, y_a, y_b, lam, criterion=\"default\"):\n        if criterion == \"default\":\n            criterion = self.criterion\n        return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n    def training_step(self, batch, batch_idx):\n        if self.cfg[\"model\"][\"train_2nd\"] and self.current_epoch >= (self.cfg[\"pl_params\"][\"max_epochs\"] - self.cfg[\"model\"][\"epoch_2nd\"]):\n            # 最後だけaugmentation切るとかする用\n            self.cfg[\"model\"][\"aug_mix\"] = False      \n        if self.cfg[\"model\"][\"aug_mix\"] and torch.rand(1) < 0.5:\n            if self.cfg[\"arcface\"] is not None:\n                X, y = batch\n                mixed_X, y_a, y_b, lam = self.mixup_data(X, y)\n                pred_y = self.forward(mixed_X)\n                pred_y = self.arc(pred_y, y, self.device)\n                loss = self.mix_criterion(pred_y ,y_a.long().squeeze(), y_b.long().squeeze(), lam)\n                self.log(\"train_loss\", loss, prog_bar=True)\n                return loss\n            elif self.cfg[\"task\"][\"aux_target\"] is not None:\n                X, y, aux_y = batch\n                mixed_X, y_a, y_b, lam, index = self.mixup_data(X, y, return_index=True)\n                pred_y, pred_aux_y = self.forward(mixed_X)\n                aux_y_a, aux_y_b = aux_y, aux_y[index]\n                cancer_loss = self.mix_criterion(pred_y ,y_a, y_b, lam)\n                aux_loss = torch.mean(torch.stack([self.mix_criterion(pred_aux_y[i], aux_y_a[:, i], aux_y_b[:, i], lam, criterion=torch.nn.functional.cross_entropy) for i in range(aux_y.shape[-1])]))\n                loss = cancer_loss + self.cfg[\"task\"][\"aux_loss_weight\"] * aux_loss\n                self.log(\"cancer_loss\", cancer_loss, prog_bar=False)\n                self.log(\"aux_loss\", aux_loss, prog_bar=False)\n                self.log(\"train_loss\", loss, prog_bar=True)\n                return {\"loss\": loss, \"cancer_loss\": cancer_loss, \"aux_loss\": aux_loss}\n            else:\n                X, y = batch\n                #if torch.rand(1) >= 0.5:\n                #    mixed_X, y_a, y_b, lam = self.mixup_data(X, y)\n                #else:\n                #    mixed_X, y_a, y_b, lam = self.cutmix_data(X, y)\n                mixed_X, y_a, y_b, lam = self.mixup_data(X, y)\n                pred_y = self.forward(mixed_X)\n                loss = self.mix_criterion(pred_y ,y_a, y_b, lam)\n                self.log(\"train_loss\", loss, prog_bar=True)\n                return loss\n        else:\n            if self.cfg[\"arcface\"] is not None:\n                X, y = batch\n                pred_y = self.arc(pred_y, y, self.device)\n                loss = self.criterion(pred_y, y.long().squeeze())\n                self.log(\"train_loss\", loss, prog_bar=True)\n                return loss\n            elif self.cfg[\"task\"][\"aux_target\"] is not None:\n                X, y, aux_y = batch\n                pred_y, pred_aux_y = self.forward(X)\n                cancer_loss = self.criterion(pred_y, y)\n                aux_loss = torch.mean(torch.stack([torch.nn.functional.cross_entropy(pred_aux_y[i], aux_y[:, i]) for i in range(aux_y.shape[-1])]))\n                loss = cancer_loss + self.cfg[\"task\"][\"aux_loss_weight\"] * aux_loss\n                self.log(\"cancer_loss\", cancer_loss, prog_bar=False)\n                self.log(\"aux_loss\", aux_loss, prog_bar=False)\n                self.log(\"train_loss\", loss, prog_bar=True)\n                return {\"loss\": loss, \"cancer_loss\": cancer_loss, \"aux_loss\": aux_loss}\n            else:\n                X, y = batch\n                pred_y = self.forward(X)\n                loss = self.criterion(pred_y, y)\n                self.log(\"train_loss\", loss, prog_bar=True)\n                return loss\n\n    def training_epoch_end(self, outputs):\n        loss_list = [x[\"loss\"] for x in outputs]\n        avg_loss = torch.stack(loss_list).mean()\n        self.log(\"train_avg_loss\", avg_loss, prog_bar=True)\n        \n        if self.cfg[\"task\"][\"aux_target\"] is not None:\n            cancer_loss_list = [x[\"cancer_loss\"] for x in outputs]\n            aux_loss_list = [x[\"aux_loss\"] for x in outputs]\n            cancer_avg_loss = torch.stack(cancer_loss_list).mean()\n            aux_avg_loss = torch.stack(aux_loss_list).mean()\n            self.log(\"train_avg_cancer_loss\", cancer_avg_loss, prog_bar=False)\n            self.log(\"train_avg_aux_loss\", aux_avg_loss, prog_bar=False)\n\n    def validation_step(self, batch, batch_idx):\n        if self.cfg[\"arcface\"] is not None:\n            X, y = batch\n            pred_y = self.forward(X)\n            pred_y = self.arc(pred_y, y, self.device)\n            loss = self.criterion(pred_y, y.long().squeeze())\n            pred_y = nn.Softmax(dim=1)(pred_y)\n            pred_y = pred_y[:, 1]\n            pred_y = torch.nan_to_num(pred_y)\n            return {\"valid_loss\": loss, \"preds\": pred_y, \"targets\": y}\n        elif self.cfg[\"task\"][\"aux_target\"] is not None:\n            X, y, aux_y = batch\n            pred_y, pred_aux_y = self.forward(X)\n            cancer_loss = self.criterion(pred_y, y)\n            aux_loss = torch.mean(torch.stack([torch.nn.functional.cross_entropy(pred_aux_y[i], aux_y[:, i]) for i in range(aux_y.shape[-1])]))\n            loss = cancer_loss + self.cfg[\"task\"][\"aux_loss_weight\"] * aux_loss\n            pred_y = torch.sigmoid(pred_y)\n            pred_y = torch.nan_to_num(pred_y)\n            return {\"valid_loss\": loss, \"cancer_loss\": cancer_loss, \"aux_loss\": aux_loss, \"preds\": pred_y, \"targets\": y}\n        else:\n            X, y = batch\n            pred_y = self.forward(X)\n            loss = self.criterion(pred_y, y)\n            pred_y = torch.sigmoid(pred_y)\n            pred_y = torch.nan_to_num(pred_y)\n            return {\"valid_loss\": loss, \"preds\": pred_y, \"targets\": y}\n\n    def validation_epoch_end(self, outputs):\n        loss_list = [x[\"valid_loss\"] for x in outputs]\n        preds = torch.cat([x[\"preds\"] for x in outputs], dim=0).cpu().detach().numpy()\n        targets = (\n            torch.cat([x[\"targets\"] for x in outputs], dim=0).cpu().detach().numpy()\n        )\n        avg_loss = torch.stack(loss_list).mean()\n        pfbeta_score = pfbeta(targets.flatten(), preds.flatten())\n        if np.unique(targets).shape[0] == 1:\n            auc_score = 0.0\n        else:\n            auc_score = sklearn.metrics.roc_auc_score(targets.flatten(), preds.flatten())\n        optimized_pfbeta_score, threshold = optimal_f1(targets.flatten(), preds.flatten())\n        recall = sklearn.metrics.recall_score(targets.flatten(), preds.flatten() > threshold)\n        specificity = sklearn.metrics.recall_score(targets.flatten(), preds.flatten() > threshold, pos_label=0)\n        precision = sklearn.metrics.precision_score(targets.flatten(), preds.flatten() > threshold)\n        self.log(\"valid_avg_loss\", avg_loss, prog_bar=True)\n        self.log(\"valid_pfbeta_score\", pfbeta_score, prog_bar=True)\n        self.log(\"optimized_pfbeta_score\", optimized_pfbeta_score, prog_bar=True)\n        self.log(\"valid_auc_score\", auc_score, prog_bar=True)\n        self.log(\"threshold\", threshold, prog_bar=False)\n        self.log(\"recall\", recall, prog_bar=False)\n        self.log(\"specificity\", specificity, prog_bar=False)\n        self.log(\"precision\", precision, prog_bar=False)\n        \n        if self.cfg[\"task\"][\"aux_target\"] is not None:\n            cancer_loss_list = [x[\"cancer_loss\"] for x in outputs]\n            aux_loss_list = [x[\"aux_loss\"] for x in outputs]\n            cancer_avg_loss = torch.stack(cancer_loss_list).mean()\n            aux_avg_loss = torch.stack(aux_loss_list).mean()\n            self.log(\"valid_avg_cancer_loss\", cancer_avg_loss, prog_bar=False)\n            self.log(\"valid_avg_aux_loss\", aux_avg_loss, prog_bar=False)\n        \n        return avg_loss\n\n    def predict_step(self, batch, batch_idx, dataloader_idx=0):\n        if self.cfg[\"tta\"]:\n            X, X_2, y = batch\n            \n            \n            if self.cfg[\"arcface\"] is not None:\n                pred_y_1 = self.forward(X)\n                pred_y_2 = self.forward(X_2)\n                pred_y_1 = self.arc(pred_y_1, y, self.device)      \n                pred_y_2 = self.arc(pred_y_2, y, self.device) \n                pred_y_1 = nn.Softmax(dim=1)(pred_y_1)\n                pred_y_2 = nn.Softmax(dim=1)(pred_y_2)\n                pred_y_1 = pred_y_1[:, 1]\n                pred_y_2 = pred_y_2[:, 1]\n            elif self.cfg[\"task\"][\"aux_target\"] is not None:\n                pred_y_1, _ = self.forward(X)\n                pred_y_2, _ = self.forward(X_2)\n                pred_y_1 = torch.sigmoid(pred_y_1)\n                pred_y_2 = torch.sigmoid(pred_y_2)\n            else:\n                pred_y_1 = self.forward(X)\n                pred_y_2 = self.forward(X_2)\n                pred_y_1 = torch.sigmoid(pred_y_1)\n                pred_y_2 = torch.sigmoid(pred_y_2)\n                \n            pred_y = (pred_y_1 + pred_y_2) / 2.0\n        else:\n            X, y = batch\n            \n            if self.cfg[\"arcface\"] is not None:\n                pred_y = self.forward(X)\n                pred_y = self.arc(pred_y, y, self.device)        \n                pred_y = nn.Softmax(dim=1)(pred_y)\n                pred_y = pred_y[:, 1]\n            elif self.cfg[\"task\"][\"aux_target\"] is not None:\n                pred_y, _ = self.forward(X)\n                pred_y = torch.sigmoid(pred_y)\n            else:\n                pred_y = self.forward(X)\n                pred_y = torch.sigmoid(pred_y)\n                \n        return pred_y\n\n    def configure_optimizers(self):\n        optimizer = optim.__dict__[self.cfg[\"model\"][\"optimizer\"][\"name\"]](\n            self.parameters(), **self.cfg[\"model\"][\"optimizer\"][\"params\"]\n        )\n        if self.cfg[\"model\"][\"scheduler\"] is None:\n            return [optimizer]\n        else:\n            if self.cfg[\"model\"][\"scheduler\"][\"name\"] == \"OneCycleLR\":\n                scheduler = optim.lr_scheduler.OneCycleLR(\n                    optimizer,\n                    #steps_per_epoch=self.cfg[\"len_train_loader\"] // self.cfg[\"pl_params\"][\"accumulate_grad_batches\"],\n                    total_steps=self.trainer.estimated_stepping_batches,\n                    **self.cfg[\"model\"][\"scheduler\"][\"params\"],\n                )\n                scheduler = {\"scheduler\": scheduler, \"interval\": \"step\"}\n            elif self.cfg[\"model\"][\"scheduler\"][\"name\"] == \"ReduceLROnPlateau\":\n                scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n                    optimizer, **self.cfg[\"model\"][\"scheduler\"][\"params\"],\n                )\n                scheduler = {\n                    \"scheduler\": scheduler,\n                    \"interval\": \"epoch\",\n                    \"monitor\": \"valid_avg_loss\",\n                }\n            else:\n                scheduler = optim.lr_scheduler.__dict__[\n                    self.cfg[\"model\"][\"scheduler\"][\"name\"]\n                ](optimizer, **self.cfg[\"model\"][\"scheduler\"][\"params\"])\n            return [optimizer], [scheduler]","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.678884Z","iopub.execute_input":"2023-02-27T10:15:35.682328Z","iopub.status.idle":"2023-02-27T10:15:35.766311Z","shell.execute_reply.started":"2023-02-27T10:15:35.682277Z","shell.execute_reply":"2023-02-27T10:15:35.765115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNADataset(torch.utils.data.Dataset):\n    def __init__(self, cfg, X, y=None, augmentation=False, cutout=False, aux=False):\n        self.cfg = cfg\n        self.augmentation = augmentation\n        self.cutout = cutout\n        self.df = X\n        self.aux = aux\n        \n        if cfg[\"model\"][\"in_chans\"] == 3:\n            self.mean = cfg[\"model\"][\"mean\"]\n            self.std = cfg[\"model\"][\"std\"]\n        else:\n            self.mean = cfg[\"model\"][\"mean\"]\n            self.std = cfg[\"model\"][\"std\"]\n        \n        if y is None:\n            self.y = torch.zeros(len(self.df), dtype=torch.float32)\n        else:\n            self.y = torch.tensor(y.values, dtype=torch.float32)\n        \n        # normalize\n        self.normalize = torchvision.transforms.Compose([\n            torchvision.transforms.ToTensor(),\n            torchvision.transforms.Normalize(mean=self.mean, std=self.std),\n            torchvision.transforms.Resize((cfg[\"task\"][\"height_size\"], cfg[\"task\"][\"width_size\"])),\n            #torchvision.transforms.Normalize(mean=self.mean, std=self.std),\n            \n            #torchvision.transforms.Resize(min(cfg[\"task\"][\"height_size\"], cfg[\"task\"][\"width_size\"])),\n            #torchvision.transforms.CenterCrop((cfg[\"task\"][\"height_size\"], cfg[\"task\"][\"width_size\"])),\n        ])\n        # resize\n        #self.resize = torchvision.transforms.CenterCrop((cfg[\"task\"][\"height_size\"], cfg[\"task\"][\"width_size\"]))\n        \n        \n        # augmentation\n        # flip\n        self.aug_hori_flip = A.HorizontalFlip(p=0.5)\n        self.aug_ver_flip = A.VerticalFlip(p=0.5)\n        # elastic and grid\n        self.aug_distortion = A.GridDistortion(p=0.5)\n        \"\"\"\n        A.OneOf([\n            A.ElasticTransform(p=0.5),\n            A.GridDistortion(p=0.5)\n        ], p=0.5)\n        \"\"\"\n        # affine\n        #self.aug_affine = A.Affine(scale=(0.8, 1.2), translate_percent=(0.0, 0.1), rotate=(-45, 45), shear=(-15, 15), p=0.8)\n        self.aug_affine = A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.2, rotate_limit=45, p=0.8)\n        # clahe\n        self.aug_clahe = A.CLAHE(p=0.5)\n        # bright\n        self.aug_bright = A.OneOf([\n            A.RandomGamma(gamma_limit=(50, 150), p=0.5),\n            A.RandomBrightnessContrast(brightness_limit=0.5, contrast_limit=0.5, p=0.5)\n        ], p=0.5)\n        # cutout\n        self.aug_cutout = A.CoarseDropout(max_height=8, max_width=8, p=0.5)\n        # randomcrop\n        #self.randomcrop = A.RandomResizedCrop(height=cfg[\"task\"][\"height_size\"], width=cfg[\"task\"][\"width_size\"], scale=(0.1, 1.0), ratio=(0.5, 1.0), p=0.8)\n\n    def __len__(self):\n        return len(self.df)\n    \n    def image_resize(self, image, width = None, height = None, inter = cv2.INTER_LINEAR):\n        dim = None\n        (h, w) = image.shape[:2]\n\n        if width is None and height is None:\n            return image\n\n        if width is None:\n            r = height / float(h)\n            dim = (int(w * r), height)\n        else:\n            r = width / float(w)\n            dim = (width, int(h * r))\n        resized = cv2.resize(image, dim, interpolation = inter)\n\n        return resized\n\n    def __getitem__(self, index):\n        image_id = self.df.loc[index, \"image_id\"]\n        patient_id = self.df.loc[index, \"patient_id\"]\n        extend = self.cfg[\"task\"][\"extend\"]\n        size = self.cfg[\"task\"][\"size\"]\n        bit = self.cfg[\"task\"][\"bit\"]\n        if self.cfg[\"task\"][\"voi\"]:\n            image_path = (\n                f\"{SAVE_FOLDER}/{patient_id}_{image_id}\"\n            )\n        else:\n            image_path = (\n                f\"{SAVE_FOLDER}/{patient_id}_{image_id}\"\n            )\n            \n        if self.cfg[\"model\"][\"in_chans\"] == 3:\n            X = cv2.imread(f\"{image_path}.png\")\n        else:\n            X = cv2.imread(f\"{image_path}.png\", cv2.IMREAD_GRAYSCALE)\n        \n        \"\"\"\n        shape = X.shape\n        if shape[0] > shape[1]:\n            X = self.image_resize(X, height=self.cfg[\"task\"][\"height_size\"])\n        else:\n            X = self.image_resize(X, width=self.cfg[\"task\"][\"width_size\"])\n        X = Image.fromarray(X)\n        X = self.resize(X)\n        X = np.array(X)\n        \"\"\"\n    \n        # augmentation\n        if self.augmentation:\n            X = self.aug_hori_flip(image=X)[\"image\"]\n            X = self.aug_ver_flip(image=X)[\"image\"]\n            #X = self.aug_distortion(image=X)[\"image\"]\n            #X = self.aug_clahe(image=X)[\"image\"]\n            X = self.aug_affine(image=X)[\"image\"]\n            X = self.aug_bright(image=X)[\"image\"]\n            if self.cutout:\n                X = self.aug_cutout(image=X)[\"image\"]\n            #X = self.randomcrop(image=X)[\"image\"]\n            \n        y = self.y[index].unsqueeze(0)\n            \n        if self.cfg[\"tta\"]:\n            X_2 = self.aug_hori_flip(image=X)[\"image\"]\n            X = self.normalize(X)\n            X_2 = self.normalize(X_2)\n            return X, X_2, y\n        else:\n            X = self.normalize(X)\n            #X = torchvision.transforms.ToTensor()(X)\n            #if self.normalize:\n            #    if self.cfg[\"model\"][\"normalize_method\"] == \"z_intensity\":\n            #        X = torchvision.transforms.Normalize(mean=torch.mean(X, dim=(1, 2)), std=torch.std(X, dim=(1, 2))+1e-7)(X)\n            #    else:\n            #        X = torchvision.transforms.Normalize(mean=self.mean, std=self.std)\n            #X = torchvision.transforms.Resize((self.cfg[\"task\"][\"height_size\"], self.cfg[\"task\"][\"width_size\"]))(X)\n            if self.aux:\n                aux_y = self.df.iloc[index][self.cfg[\"task\"][\"aux_target\"]]\n                aux_y = torch.tensor(aux_y.values, dtype=torch.long)\n                return X, y, aux_y\n            else:\n                return X, y","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.769334Z","iopub.execute_input":"2023-02-27T10:15:35.770371Z","iopub.status.idle":"2023-02-27T10:15:35.793659Z","shell.execute_reply.started":"2023-02-27T10:15:35.770268Z","shell.execute_reply":"2023-02-27T10:15:35.792399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CNNClassifierInference:\n    def __init__(self, cfg, weight_path=None):\n        # aux\n        if cfg[\"task\"][\"aux_target\"] is not None:\n            cfg[\"task\"][\"aux_target_nclasses\"] = []\n            for t in cfg[\"task\"][\"aux_target\"]:\n                if t == \"site_id\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(2)\n                elif t == \"laterality\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(2)\n                elif t == \"view\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(6)\n                elif t == \"implant\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(2)\n                elif t == \"biopsy\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(2)\n                elif t == \"invasive\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(2)\n                elif t == \"BIRADS\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(4)\n                elif t == \"density\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(5)\n                elif t == \"difficult_negative_case\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(2)\n                elif t == \"age\":\n                    cfg[\"task\"][\"aux_target_nclasses\"].append(10)\n            print(cfg[\"task\"][\"aux_target_nclasses\"])\n        else:\n            aux = False\n            \n        self.weight_path = weight_path\n        self.cfg = cfg\n        if cfg[\"model\"][\"weighted_loss\"]:\n            self.cfg[\"model\"][\"loss_weights\"] = torch.tensor(0.0, dtype=torch.float32)\n        else:\n            self.cfg[\"model\"][\"loss_weights\"] = None\n        self.model = CNNModel(self.cfg)\n        self.trainer = Trainer(**self.cfg[\"pl_params\"])\n\n    def predict(self, test_X):\n        test_dataset = RSNADataset(self.cfg, test_X)\n        test_dataloader = torch.utils.data.DataLoader(\n            test_dataset, **self.cfg[\"test_loader\"],\n        )\n        preds = self.trainer.predict(\n            self.model, dataloaders=test_dataloader, ckpt_path=self.weight_path\n        )\n        preds = torch.cat(preds, axis=0)\n        preds = preds.cpu().detach().numpy()\n\n        return preds","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.795829Z","iopub.execute_input":"2023-02-27T10:15:35.796129Z","iopub.status.idle":"2023-02-27T10:15:35.812017Z","shell.execute_reply.started":"2023-02-27T10:15:35.7961Z","shell.execute_reply":"2023-02-27T10:15:35.81068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def one_fold_test(cfg, test_X, fold_n):\n    print(f\"[fold_{fold_n}]\")\n    seed_everything(cfg[\"general\"][\"seed\"], workers=True)\n\n    model = CNNClassifierInference(cfg, f\"{cfg['ckpt_path']}.ckpt\")\n    test_preds = model.predict(test_X)\n\n    del model\n    gc.collect()\n    torch.cuda.empty_cache()\n\n    return test_preds","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.813927Z","iopub.execute_input":"2023-02-27T10:15:35.814358Z","iopub.status.idle":"2023-02-27T10:15:35.82592Z","shell.execute_reply.started":"2023-02-27T10:15:35.814317Z","shell.execute_reply":"2023-02-27T10:15:35.825141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def all_test(cfg, test_X):\n    print(f\"[all]\")\n    seed_everything(cfg[\"general\"][\"seed\"], workers=True)\n\n    model = CNNClassifierInference(cfg, f\"{cfg['ckpt_path']}.ckpt\")\n    test_preds = model.predict(test_X)\n\n    del model\n    gc.collect()\n    torch.cuda.empty_cache()\n\n    return test_preds","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.8276Z","iopub.execute_input":"2023-02-27T10:15:35.828724Z","iopub.status.idle":"2023-02-27T10:15:35.84026Z","shell.execute_reply.started":"2023-02-27T10:15:35.828687Z","shell.execute_reply":"2023-02-27T10:15:35.839244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# main","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(f\"{cfg1['general']['input_path']}/train.csv\")\ntrain_X = train.drop(cfg1[\"task\"][\"target\"], axis=1)\n\ntest_X = pd.read_csv(f\"{cfg1['general']['input_path']}/test.csv\")\nsub = pd.read_csv(f\"{cfg1['general']['input_path']}/sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.842009Z","iopub.execute_input":"2023-02-27T10:15:35.842409Z","iopub.status.idle":"2023-02-27T10:15:35.957999Z","shell.execute_reply.started":"2023-02-27T10:15:35.842353Z","shell.execute_reply":"2023-02-27T10:15:35.956792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if USE_DATA == \"test\":\n    X = test_X\nelse:\n    X = train_X","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.95971Z","iopub.execute_input":"2023-02-27T10:15:35.960356Z","iopub.status.idle":"2023-02-27T10:15:35.965903Z","shell.execute_reply.started":"2023-02-27T10:15:35.960316Z","shell.execute_reply":"2023-02-27T10:15:35.964707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# random seed setting\nseed_everything(cfg1[\"general\"][\"seed\"], workers=True)\n\nfinal_test_preds = []\nif cfg1[\"general\"][\"cv\"]:\n    test_preds_list = []\n    for fold_n in tqdm(cfg1[\"general\"][\"fold\"]):\n        for cfg in cfg_list:\n            print(cfg[\"general\"][\"save_name\"])\n            cfg[\"fold_n\"] = fold_n\n            cfg[\"ckpt_path\"] = f\"{cfg['general']['output_path']}/{cfg['general']['save_name']}/last_epoch_fold{fold_n}\"\n            test_preds = one_fold_test(cfg, X, fold_n)\n            test_preds_list.append(test_preds)\n            print(test_preds)\n            del test_preds\n        final_test_preds.append(np.mean(test_preds_list, axis=0))\n        print(\"-----------------\")\n        print(\"ensemble: \")\n        print(np.mean(test_preds_list, axis=0))\n        print(\"-----------------\")\n    final_test_preds = np.mean(final_test_preds, axis=0)\n    print()\n    print(\"final_test_preds:\")\n    print(final_test_preds)\nelse:\n    test_preds_list = []\n    for cfg in cfg_list:\n        cfg[\"fold_n\"] = \"all\"\n        cfg[\"ckpt_path\"] = f\"{cfg['general']['output_path']}/{cfg['general']['save_name']}/last_epoch\"\n        test_preds = all_test(cfg, test_X)\n        test_preds_list.append(test_preds)\n        print(test_preds)\n        del test_preds\n    final_test_preds = np.mean(test_preds_list, axis=0)\n    print()\n    print(\"final_test_preds:\")\n    print(final_test_preds)","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:35.974371Z","iopub.execute_input":"2023-02-27T10:15:35.975157Z","iopub.status.idle":"2023-02-27T10:15:55.044096Z","shell.execute_reply.started":"2023-02-27T10:15:35.975116Z","shell.execute_reply":"2023-02-27T10:15:55.0429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = X.copy()\nsub[\"cancer\"] = final_test_preds\n\nif AGG == \"max\":\n    sub = sub[[\"prediction_id\", \"cancer\"]].groupby(\"prediction_id\").max().reset_index()\nelif AGG == \"mean\":\n    sub = sub[[\"prediction_id\", \"cancer\"]].groupby(\"prediction_id\").mean().reset_index()\nelif AGG == \"top3_mean\":\n    sub = sub[[\"prediction_id\", \"cancer\"]].groupby(\"prediction_id\").head(3).groupby(\"prediction_id\").mean().reset_index()","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:55.045927Z","iopub.execute_input":"2023-02-27T10:15:55.046619Z","iopub.status.idle":"2023-02-27T10:15:55.065979Z","shell.execute_reply.started":"2023-02-27T10:15:55.046575Z","shell.execute_reply":"2023-02-27T10:15:55.064843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:55.069184Z","iopub.execute_input":"2023-02-27T10:15:55.069527Z","iopub.status.idle":"2023-02-27T10:15:55.086528Z","shell.execute_reply.started":"2023-02-27T10:15:55.069498Z","shell.execute_reply":"2023-02-27T10:15:55.085447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if BINARIZED:\n    sub[\"cancer\"] = sub[\"cancer\"] > THR\n    sub = sub.astype({\"cancer\": float})","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:55.088207Z","iopub.execute_input":"2023-02-27T10:15:55.088551Z","iopub.status.idle":"2023-02-27T10:15:55.09512Z","shell.execute_reply.started":"2023-02-27T10:15:55.088515Z","shell.execute_reply":"2023-02-27T10:15:55.093993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:55.096753Z","iopub.execute_input":"2023-02-27T10:15:55.097455Z","iopub.status.idle":"2023-02-27T10:15:55.110819Z","shell.execute_reply.started":"2023-02-27T10:15:55.097415Z","shell.execute_reply":"2023-02-27T10:15:55.109406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.to_csv('/kaggle/working/submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-02-27T10:15:55.112551Z","iopub.execute_input":"2023-02-27T10:15:55.113066Z","iopub.status.idle":"2023-02-27T10:15:55.1235Z","shell.execute_reply.started":"2023-02-27T10:15:55.113027Z","shell.execute_reply":"2023-02-27T10:15:55.122444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}