{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52254,"databundleVersionId":6863140,"sourceType":"competition"},{"sourceId":6205800,"sourceType":"datasetVersion","datasetId":3563333},{"sourceId":6211844,"sourceType":"datasetVersion","datasetId":3567114}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Install Libraries","metadata":{}},{"cell_type":"code","source":"# !pip install -q keras-cv-attention-models\n!pip install -qU wandb\n!pip install -qU scikit-learn\n!pip install -q seaborn\n! pip install torchgeometry","metadata":{"_kg_hide-output":true,"_kg_hide-input":false,"execution":{"iopub.status.busy":"2024-03-11T16:34:01.881018Z","iopub.execute_input":"2024-03-11T16:34:01.881671Z","iopub.status.idle":"2024-03-11T16:34:59.917386Z","shell.execute_reply.started":"2024-03-11T16:34:01.881639Z","shell.execute_reply":"2024-03-11T16:34:59.916156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import Libraries","metadata":{}},{"cell_type":"code","source":"import os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'  # to avoid too many logging messages\nimport pandas as pd, numpy as np, random, shutil\nimport tensorflow as tf, re, math\nimport tensorflow.keras.backend as K\nimport sklearn\nimport matplotlib.pyplot as plt\n# import tensorflow_addons as tfa\nimport tensorflow_probability as tfp\nimport wandb\nimport yaml\nimport wandb\nfrom IPython import display as ipd\nfrom glob import glob\nfrom tqdm import tqdm\nfrom sklearn.model_selection import KFold, StratifiedKFold, GroupKFold, StratifiedGroupKFold\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.utils.class_weight import compute_class_weight\nimport torch\nimport torchvision.transforms as transforms\nimport torch.nn.functional as F\nimport torchgeometry as tgm\nimport math","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-03-11T16:34:59.919224Z","iopub.execute_input":"2024-03-11T16:34:59.919528Z","iopub.status.idle":"2024-03-11T16:35:19.153774Z","shell.execute_reply.started":"2024-03-11T16:34:59.919503Z","shell.execute_reply":"2024-03-11T16:35:19.152749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key = user_secrets.get_secret(\"WANDB\")\n\n    wandb.login(key=api_key)\n    anonymous = None\nexcept:\n    anonymous = \"must\"\n    print('To use your W&B account,\\nGo to Add-ons -> Secrets and provide your W&B access token. Use the Label name as WANDB. \\nGet your W&B access token from here: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:19.155617Z","iopub.execute_input":"2024-03-11T16:35:19.156361Z","iopub.status.idle":"2024-03-11T16:35:19.326592Z","shell.execute_reply.started":"2024-03-11T16:35:19.156324Z","shell.execute_reply":"2024-03-11T16:35:19.325635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    wandb         = True\n    competition   = 'rsna-atd' \n    _wandb_kernel = 'awsaf49'\n    debug         = False\n    comment       = 'EfficientNetV1B0-256x256-low_lr-vflip'\n    exp_name      = 'baseline-v4: new_ds + multi_head' # name of the experiment, folds will be grouped using 'exp_name'\n    verbose      = 0\n    display_plot = True\n    device = \"GPU\" #or \"GPU\"\n    model_name = 'EfficientNetV1B0'\n    seed = 42\n    folds = 4\n    selected_folds = [0, 1, 2]\n    img_size = [256, 256]\n    batch_size = 24\n    epochs = 10\n    loss      = 'BCE & CCE'  # BCE, Focal\n    optimizer = 'Adam'\n    augment   = True\n    transform = 0.90  # transform prob\n    fill_mode = 'constant'\n    rot    = 2.0\n    shr    = 2.0\n    hzoom  = 50.0\n    wzoom  = 50.0\n    hshift = 10.0\n    wshift = 10.0\n    # flip\n    hflip = True\n    vflip = True\n    # clip\n    clip = False\n    # lr-scheduler\n    scheduler   = 'cosine' # cosine\n    # dropout\n    drop_prob   = 0.6\n    drop_cnt    = 5\n    drop_size   = 0.05\n    # cut-mix-up\n    mixup_prob = 0.0\n    mixup_alpha = 0.5\n    cutmix_prob = 0.0\n    cutmix_alpha = 2.5\n    # pixel-augment\n    pixel_aug = 0.90  # prob of pixel_aug\n    sat  = [0.7, 1.3]\n    cont = [0.8, 1.2]\n    bri  = 0.15\n    hue  = 0.05\n    # test-time augs\n    tta = 1\n    # target column\n    target_col  = [ \"bowel_injury\", \"extravasation_injury\", \"kidney_healthy\", \"kidney_low\",\n                   \"kidney_high\", \"liver_healthy\", \"liver_low\", \"liver_high\",\n                   \"spleen_healthy\", \"spleen_low\", \"spleen_high\"] # not using \"bowel_healthy\" & \"extravasation_healthy\"\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:19.328902Z","iopub.execute_input":"2024-03-11T16:35:19.329212Z","iopub.status.idle":"2024-03-11T16:35:19.339115Z","shell.execute_reply.started":"2024-03-11T16:35:19.329185Z","shell.execute_reply":"2024-03-11T16:35:19.338036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reproducibility","metadata":{}},{"cell_type":"code","source":"def seeding(SEED):\n    np.random.seed(SEED)\n    random.seed(SEED)\n    os.environ['PYTHONHASHSEED'] = str(SEED)\n    tf.random.set_seed(SEED)\n    print('seeding done!!!')\nseeding(CFG.seed)\ndevice = \"cuda\" if CFG.device == 'GPU' else \"cpu\"\nBASE_PATH = f'/kaggle/input/rsna-atd-512x512-png-v2-dataset'","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:19.340258Z","iopub.execute_input":"2024-03-11T16:35:19.340578Z","iopub.status.idle":"2024-03-11T16:35:19.353873Z","shell.execute_reply.started":"2024-03-11T16:35:19.340553Z","shell.execute_reply":"2024-03-11T16:35:19.352984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train\ndf = pd.read_csv(f'{BASE_PATH}/train.csv')\ndf['image_path'] = f'{BASE_PATH}/train_images'\\\n                    + '/' + df.patient_id.astype(str)\\\n                    + '/' + df.series_id.astype(str)\\\n                    + '/' + df.instance_number.astype(str) +'.png'\ndf = df.drop_duplicates()\nprint('Train:')\n# display(df.head(2))\n\n# test\ntest_df = pd.read_csv(f'{BASE_PATH}/test.csv')\ntest_df['image_path'] = f'{BASE_PATH}/test_images'\\\n                    + '/' + test_df.patient_id.astype(str)\\\n                    + '/' + test_df.series_id.astype(str)\\\n                    + '/' + test_df.instance_number.astype(str) +'.png'\ntest_df = test_df.drop_duplicates()\nprint('\\nTest:')\n# display(test_df.head(2))","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:19.354989Z","iopub.execute_input":"2024-03-11T16:35:19.355227Z","iopub.status.idle":"2024-03-11T16:35:19.488588Z","shell.execute_reply.started":"2024-03-11T16:35:19.355207Z","shell.execute_reply":"2024-03-11T16:35:19.487538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nif os.path.exists(df.image_path.iloc[0]):\n    print(f\"File exists: {df.image_path.iloc[0]}\")\nelse:\n    print(f\"File does not exist: {df.image_path.iloc[0]}\")\n\nif os.path.exists(test_df.image_path.iloc[0]):\n    print(f\"File exists: {test_df.image_path.iloc[0]}\")\nelse:\n    print(f\"File does not exist: {test_df.image_path.iloc[0]}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:19.48988Z","iopub.execute_input":"2024-03-11T16:35:19.490185Z","iopub.status.idle":"2024-03-11T16:35:19.508416Z","shell.execute_reply.started":"2024-03-11T16:35:19.49016Z","shell.execute_reply":"2024-03-11T16:35:19.507531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train-Test Ditribution","metadata":{}},{"cell_type":"code","source":"import seaborn as sns\nimport matplotlib.pyplot as plt\n\nsns.set_theme(style=\"white\")\n\n# Compute the correlation matrix\ncorr = df[df.columns[1:14]].corr().round(2)\n\n# Generate a mask for the upper triangle\nmask = np.triu(np.ones_like(corr, dtype=bool))\n\n# Generate a custom diverging colormap\ncmap = sns.diverging_palette(230, 20, as_cmap=True)\n\n# Set up the matplotlib figure\nf, ax = plt.subplots(figsize=(11, 9))\n\n# Draw the heatmap with the mask and correct aspect ratio\nsns.heatmap(corr, mask=mask, annot=True, cmap=cmap, vmax=.3, center=0,\n            square=True, linewidths=.5, cbar_kws={\"shrink\": .5})\nplt.show();","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-03-11T16:35:19.509651Z","iopub.execute_input":"2024-03-11T16:35:19.51008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Split\n* Data is splited while stratifying, 'laterality', 'view', 'age','cancer', 'biopsy', 'invasive', 'BIRADS', 'implant', 'density','machine_id', 'difficult_negative_case','cancer'.\n* To avoid leakage, data is also split keeping images from same patient in either train or valid not in both.\n* `StratifiedGroupKFold` does the both job **stratifying** and **group** split.\n\n","metadata":{}},{"cell_type":"code","source":"df['stratify'] = ''\nfor col in CFG.target_col:\n    df['stratify'] += df[col].astype(str)\n\ndf = df.reset_index(drop=True)\nskf = StratifiedGroupKFold(n_splits=CFG.folds, shuffle=True, random_state=CFG.seed)\nfor fold, (train_idx, val_idx) in enumerate(skf.split(df, df['stratify'], df[\"patient_id\"])):\n    df.loc[val_idx, 'fold'] = fold\ndisplay(df.groupby(['fold', 'patient_id']).size())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## torch版本\ndef get_mat(shear, height_zoom, width_zoom, height_shift, width_shift):\n    shear = math.pi * shear / 180.\n\n    def get_3x3_mat(lst):\n        return torch.reshape(torch.cat(lst, dim=0), [3, 3])\n\n    one = torch.tensor([1.0], dtype=torch.float32)\n    zero = torch.tensor([0.0], dtype=torch.float32)\n\n    # SHEAR MATRIX\n    c2 = torch.cos(shear)\n    print(c2)\n    s2 = torch.sin(shear)\n    print(s2)\n    shear_matrix = get_3x3_mat([one, s2, zero,\n                                zero, c2, zero,\n                                zero, zero, one])\n\n    # ZOOM MATRIX\n    zoom_matrix = get_3x3_mat([one / height_zoom, zero, zero,\n                               zero, one / width_zoom, zero,\n                               zero, zero, one])\n\n    # SHIFT MATRIX\n    shift_matrix = get_3x3_mat([one, zero, height_shift,\n                                zero, one, width_shift,\n                                zero, zero, one])\n\n    return torch.mm(shear_matrix, torch.mm(zoom_matrix, shift_matrix))\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.604535Z","iopub.execute_input":"2024-03-11T16:35:20.604856Z","iopub.status.idle":"2024-03-11T16:35:20.614806Z","shell.execute_reply.started":"2024-03-11T16:35:20.60483Z","shell.execute_reply":"2024-03-11T16:35:20.613743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"shr = CFG.shr * torch.randn(1, dtype=torch.float32)\nh_zoom = 1.0 + torch.randn(1, dtype=torch.float32) / CFG.hzoom\nw_zoom = 1.0 + torch.randn(1, dtype=torch.float32) / CFG.wzoom\nh_shift = CFG.hshift * torch.randn(1, dtype=torch.float32)\nw_shift = CFG.wshift * torch.randn(1, dtype=torch.float32)\nprint(shr,h_zoom,w_zoom,h_shift,w_shift)\nget_mat_torch(shr,h_zoom,w_zoom,h_shift,w_shift)","metadata":{"execution":{"iopub.status.busy":"2024-03-11T08:57:40.384451Z","iopub.execute_input":"2024-03-11T08:57:40.384854Z","iopub.status.idle":"2024-03-11T08:57:40.400033Z","shell.execute_reply.started":"2024-03-11T08:57:40.384824Z","shell.execute_reply":"2024-03-11T08:57:40.398794Z"}}},{"cell_type":"markdown","source":"## 函数2","metadata":{}},{"cell_type":"markdown","source":"\ndef transform(image, DIM=CFG.img_size):#[rot,shr,h_zoom,w_zoom,h_shift,w_shift]):\n    if DIM[0]>DIM[1]:\n        diff  = (DIM[0]-DIM[1])\n        pad   = [diff//2, diff//2 + diff%2]\n        image = tf.pad(image, [[0, 0], [pad[0], pad[1]],[0, 0]])\n        NEW_DIM = DIM[0]\n    elif DIM[0]<DIM[1]:\n        diff  = (DIM[1]-DIM[0])\n        pad   = [diff//2, diff//2 + diff%2]\n        image = tf.pad(image, [[pad[0], pad[1]], [0, 0],[0, 0]])\n        NEW_DIM = DIM[1]\n    \n    rot = CFG.rot * tf.random.normal([1], dtype='float32')\n    shr = CFG.shr * tf.random.normal([1], dtype='float32') \n    h_zoom = 1.0 + tf.random.normal([1], dtype='float32') / CFG.hzoom\n    w_zoom = 1.0 + tf.random.normal([1], dtype='float32') / CFG.wzoom\n    h_shift = CFG.hshift * tf.random.normal([1], dtype='float32') \n    w_shift = CFG.wshift * tf.random.normal([1], dtype='float32') \n    \n    transformation_matrix=tf.linalg.inv(get_mat(shr,h_zoom,w_zoom,h_shift,w_shift))\n    \n    flat_tensor=tfa.image.transform_ops.matrices_to_flat_transforms(transformation_matrix)\n    \n    image=tfa.image.transform(image,flat_tensor, fill_mode=CFG.fill_mode)\n    \n    rotation = math.pi * rot / 180.\n    \n    image=tfa.image.rotate(image,-rotation, fill_mode=CFG.fill_mode)\n    \n    if DIM[0]>DIM[1]:\n        image=tf.reshape(image, [NEW_DIM, NEW_DIM,3])\n        image = image[:, pad[0]:-pad[1],:]\n    elif DIM[1]>DIM[0]:\n        image=tf.reshape(image, [NEW_DIM, NEW_DIM,3])\n        image = image[pad[0]:-pad[1],:,:]\n    image = tf.reshape(image, [*DIM, 3])    \n    return image\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T13:43:00.123988Z","iopub.execute_input":"2024-03-11T13:43:00.124625Z","iopub.status.idle":"2024-03-11T13:43:00.138939Z","shell.execute_reply.started":"2024-03-11T13:43:00.12459Z","shell.execute_reply":"2024-03-11T13:43:00.137853Z"}}},{"cell_type":"code","source":"\n\ndef transform(image, DIM=CFG.img_size):\n    if DIM[0] > DIM[1]:\n        diff = (DIM[0] - DIM[1])\n        pad = [diff // 2, diff // 2 + diff % 2]\n        image = F.pad(image, (pad[0], pad[1], 0, 0, 0, 0))\n        NEW_DIM = DIM[0]\n    elif DIM[0] < DIM[1]:\n        diff = (DIM[1] - DIM[0])\n        pad = [diff // 2, diff // 2 + diff % 2]\n        image = F.pad(image, (0, 0, pad[0], pad[1], 0, 0))\n        NEW_DIM = DIM[1]\n\n    rot = CFG.rot * torch.randn(1, dtype=torch.float32)\n    shr = CFG.shr * torch.randn(1, dtype=torch.float32)\n    h_zoom = 1.0 + torch.randn(1, dtype=torch.float32) / CFG.hzoom\n    w_zoom = 1.0 + torch.randn(1, dtype=torch.float32) / CFG.wzoom\n    h_shift = CFG.hshift * torch.randn(1, dtype=torch.float32)\n    w_shift = CFG.wshift * torch.randn(1, dtype=torch.float32)\n\n    shear_matrix = get_mat_torch(shr, h_zoom, w_zoom, h_shift, w_shift)\n    flat_tensor = shear_matrix.view(-1)\n    \n    image = F.grid_sample(image.unsqueeze(0), flat_tensor.unsqueeze(0).unsqueeze(0), align_corners=False)\n    image = image.squeeze(0)\n    \n    rotation = math.pi * rot / 180.\n    image = tgm.rotate(image.unsqueeze(0), rotation, align_corners=False)\n    image = image.squeeze(0)\n\n    if DIM[0] > DIM[1]:\n        image = image[:, pad[0]:-pad[1], :]\n    elif DIM[1] > DIM[0]:\n        image = image[pad[0]:-pad[1], :, :]\n    image = image.permute(1, 2, 0)\n    \n    return image\n\n# Note: Make sure to install the torchvision and torchgeometry packages if you haven't already:\n# pip install torchvision torchgeometry\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.615988Z","iopub.execute_input":"2024-03-11T16:35:20.616309Z","iopub.status.idle":"2024-03-11T16:35:20.629993Z","shell.execute_reply.started":"2024-03-11T16:35:20.616263Z","shell.execute_reply":"2024-03-11T16:35:20.629186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\ndef dropout(image,DIM=CFG.img_size, PROBABILITY = 0.6, CT = 5, SZ = 0.1):\n    # input image - is one image of size [dim,dim,3] not a batch of [b,dim,dim,3]\n    # output - image with CT squares of side size SZ*DIM removed\n    \n    # DO DROPOUT WITH PROBABILITY DEFINED ABOVE\n    P = tf.cast( tf.random.uniform([],0,1)<PROBABILITY, tf.int32)\n    if (P==0)|(CT==0)|(SZ==0): \n        return image\n    \n    for k in range(CT):\n        # CHOOSE RANDOM LOCATION\n        x = tf.cast( tf.random.uniform([],0,DIM[1]),tf.int32)\n        y = tf.cast( tf.random.uniform([],0,DIM[0]),tf.int32)\n        # COMPUTE SQUARE \n        WIDTH = tf.cast( SZ*min(DIM),tf.int32) * P\n        ya = tf.math.maximum(0,y-WIDTH//2)\n        yb = tf.math.minimum(DIM[0],y+WIDTH//2)\n        xa = tf.math.maximum(0,x-WIDTH//2)\n        xb = tf.math.minimum(DIM[1],x+WIDTH//2)\n        # DROPOUT IMAGE\n        one = image[ya:yb,0:xa,:]\n        two = tf.zeros([yb-ya,xb-xa,3], dtype = image.dtype) \n        three = image[ya:yb,xb:DIM[1],:]\n        middle = tf.concat([one,two,three],axis=1)\n        image = tf.concat([image[0:ya,:,:],middle,image[yb:DIM[0],:,:]],axis=0)\n        image = tf.reshape(image,[*DIM,3])\n\n#     image = tf.reshape(image,[*DIM,3])\n    return image","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-03-11T13:43:23.842045Z","iopub.execute_input":"2024-03-11T13:43:23.842874Z","iopub.status.idle":"2024-03-11T13:43:23.85457Z","shell.execute_reply.started":"2024-03-11T13:43:23.842829Z","shell.execute_reply":"2024-03-11T13:43:23.853618Z"}}},{"cell_type":"code","source":"import torch\n\ndef dropout(image, DIM=CFG.img_size, PROBABILITY=0.6, CT=5, SZ=0.1):\n    # input image - is one image of size [dim,dim,3] not a batch of [b,dim,dim,3]\n    # output - image with CT squares of side size SZ*DIM removed\n\n    # DO DROPOUT WITH PROBABILITY DEFINED ABOVE\n    P = torch.randint(0, 2, (1,), dtype=torch.int32).item()\n    if (P == 0) or (CT == 0) or (SZ == 0):\n        return image\n\n    for k in range(CT):\n        # CHOOSE RANDOM LOCATION\n        x = torch.randint(0, DIM[1], (1,), dtype=torch.int32).item()\n        y = torch.randint(0, DIM[0], (1,), dtype=torch.int32).item()\n        # COMPUTE SQUARE\n        WIDTH = int(SZ * min(DIM))\n        ya = max(0, y - WIDTH // 2)\n        yb = min(DIM[0], y + WIDTH // 2)\n        xa = max(0, x - WIDTH // 2)\n        xb = min(DIM[1], x + WIDTH // 2)\n        # DROPOUT IMAGE\n        one = image[ya:yb, 0:xa, :]\n        two = torch.zeros((yb - ya, xb - xa, 3), dtype=image.dtype)\n        three = image[ya:yb, xb:DIM[1], :]\n        middle = torch.cat([one, two, three], dim=1)\n        image = torch.cat([image[0:ya, :, :], middle, image[yb:DIM[0], :, :]], dim=0)\n        image = image.view(*DIM, 3)\n\n    return image\n\n# Note: Make sure to have the appropriate PyTorch library version installed.\n# You can install it using: pip install torch\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.631203Z","iopub.execute_input":"2024-03-11T16:35:20.631611Z","iopub.status.idle":"2024-03-11T16:35:20.644506Z","shell.execute_reply.started":"2024-03-11T16:35:20.63158Z","shell.execute_reply":"2024-03-11T16:35:20.643688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CutMixUp","metadata":{}},{"cell_type":"code","source":"# def random_int(shape=[], minval=0, maxval=1):\n#     return tf.random.uniform(\n#         shape=shape, minval=minval, maxval=maxval, dtype=tf.int32)\n\ndef random_int(shape=[], minval=0, maxval=1):\n    return torch.randint(low=minval, high=maxval, size=shape, dtype=torch.int32)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.645611Z","iopub.execute_input":"2024-03-11T16:35:20.645873Z","iopub.status.idle":"2024-03-11T16:35:20.658245Z","shell.execute_reply.started":"2024-03-11T16:35:20.645851Z","shell.execute_reply":"2024-03-11T16:35:20.65745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def random_float(shape=[], minval=0.0, maxval=1.0):\n#     rnd = tf.random.uniform(\n#         shape=shape, minval=minval, maxval=maxval, dtype=tf.float32)\n#     return rnd\ndef random_float(shape=[], minval=0.0, maxval=1.0):\n    rnd = torch.rand(size=shape, dtype=torch.float32)\n    return rnd * (maxval - minval) + minval\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.659399Z","iopub.execute_input":"2024-03-11T16:35:20.659684Z","iopub.status.idle":"2024-03-11T16:35:20.668419Z","shell.execute_reply.started":"2024-03-11T16:35:20.659661Z","shell.execute_reply":"2024-03-11T16:35:20.667623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# mixup\n# def get_mixup(alpha=0.2, prob=0.5):\n#     @tf.function\n#     def mixup(images, labels, alpha=alpha, prob=prob):\n#         if random_float() > prob:\n#             return images, labels\n\n#         image_shape = tf.shape(images)\n#         label_shape = tf.shape(labels)\n\n#         beta = tfp.distributions.Beta(alpha, alpha)\n#         lam = beta.sample(1)[0]\n\n#         images = lam * images + (1.0 - lam) * tf.roll(images, shift=1, axis=0)\n#         labels = lam * labels + (1.0 - lam) * tf.roll(labels, shift=1, axis=0)\n\n#         images = tf.reshape(images, image_shape)\n#         labels = tf.reshape(labels, label_shape)\n#         return images, labels\n#     return mixup\n# \ndef get_mixup(alpha=0.2, prob=0.5):\n    def mixup(images, labels, alpha=alpha, prob=prob):\n        if random_float_tensor() > prob:\n            return images, labels\n\n        image_shape = images.shape\n        label_shape = labels.shape\n\n        beta = dist.Beta(alpha, alpha)\n        lam = beta.sample((1,))[0]\n\n        images = lam * images + (1.0 - lam) * torch.roll(images, shifts=1, dims=0)\n        labels = lam * labels + (1.0 - lam) * torch.roll(labels, shifts=1, dims=0)\n\n        images = images.view(image_shape)\n        labels = labels.view(label_shape)\n        return images, labels\n    \n    return mixup","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.669445Z","iopub.execute_input":"2024-03-11T16:35:20.669766Z","iopub.status.idle":"2024-03-11T16:35:20.679436Z","shell.execute_reply.started":"2024-03-11T16:35:20.669741Z","shell.execute_reply":"2024-03-11T16:35:20.678658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# # cutmix\n# def get_cutmix(alpha, prob=0.5):\n#     @tf.function\n#     def cutmix(images, labels, alpha=alpha, prob=prob):\n#         if random_float() > prob:\n#             return images, labels\n#         image_shape = tf.shape(images)\n#         label_shape = tf.shape(labels)\n        \n#         W = tf.cast(image_shape[2], tf.int32)\n#         H = tf.cast(image_shape[1], tf.int32)\n\n#         beta = tfp.distributions.Beta(alpha, alpha)\n#         lam = beta.sample(1)[0]\n\n#         images_rolled = tf.roll(images, shift=1, axis=0)\n#         labels_rolled = tf.roll(labels, shift=1, axis=0)\n\n#         r_x = random_int([], minval=0, maxval=W)\n#         r_y = random_int([], minval=0, maxval=H)\n#         r = 0.5 * tf.math.sqrt(1.0 - lam)\n#         r_w_half = tf.cast(r * tf.cast(W, tf.float32), tf.int32)\n#         r_h_half = tf.cast(r * tf.cast(H, tf.float32), tf.int32)\n\n#         x1 = tf.cast(tf.clip_by_value(r_x - r_w_half, 0, W), tf.int32)\n#         x2 = tf.cast(tf.clip_by_value(r_x + r_w_half, 0, W), tf.int32)\n#         y1 = tf.cast(tf.clip_by_value(r_y - r_h_half, 0, H), tf.int32)\n#         y2 = tf.cast(tf.clip_by_value(r_y + r_h_half, 0, H), tf.int32)\n\n#         # outer-pad patch -> [0, 0, 1, 1, 0, 0]\n#         patch1 = images[:, y1:y2, x1:x2, :]  # [batch, height, width, channel]\n#         patch1 = tf.pad(\n#             patch1, [[0, 0], [y1, H - y2], [x1, W - x2], [0, 0]])  # outer-pad\n\n#         # inner-pad patch -> [1, 1, 0, 0, 1, 1]\n#         patch2 = images_rolled[:, y1:y2, x1:x2, :]\n#         patch2 = tf.pad(\n#             patch2, [[0, 0], [y1, H - y2], [x1, W - x2], [0, 0]])  # outer-pad\n#         patch2 = images_rolled - patch2  # inner-pad = img - outer-pad\n\n#         images = patch1 + patch2  # cutmix img\n\n#         lam = tf.cast((1.0 - (x2 - x1) * (y2 - y1) / (W * H)), tf.float32)  # no H as (y1 - y2)/H = 1\n#         labels = lam * labels + (1.0 - lam) * labels_rolled  # cutmix label\n\n#         images = tf.reshape(images, image_shape)\n#         labels = tf.reshape(labels, label_shape)\n\n#         return images, labels\n\n#     return cutmix\ndef get_cutmix(alpha, prob=0.5):\n    def cutmix(images, labels, alpha=alpha, prob=prob):\n        if random_float_tensor() > prob:\n            return images, labels\n\n        image_shape = images.shape\n        label_shape = labels.shape\n\n        W = image_shape[3]\n        H = image_shape[2]\n\n        beta = dist.Beta(alpha, alpha)\n        lam = beta.sample((1,))[0]\n\n        images_rolled = torch.roll(images, shifts=1, dims=0)\n        labels_rolled = torch.roll(labels, shifts=1, dims=0)\n\n        r_x = random_int_tensor([], minval=0, maxval=W)\n        r_y = random_int_tensor([], minval=0, maxval=H)\n        r = 0.5 * torch.sqrt(1.0 - lam)\n        r_w_half = int(r * float(W))\n        r_h_half = int(r * float(H))\n\n        x1 = max(int(torch.clamp(r_x - r_w_half, 0, W)), 0)\n        x2 = min(int(torch.clamp(r_x + r_w_half, 0, W)), W)\n        y1 = max(int(torch.clamp(r_y - r_h_half, 0, H)), 0)\n        y2 = min(int(torch.clamp(r_y + r_h_half, 0, H)), H)\n\n        # outer-pad patch -> [0, 0, 1, 1, 0, 0]\n        patch1 = images[:, :, y1:y2, x1:x2]  # [batch, channel, height, width]\n        patch1 = F.pad(patch1, (x1, W - x2, y1, H - y2))  # outer-pad\n\n        # inner-pad patch -> [1, 1, 0, 0, 1, 1]\n        patch2 = images_rolled[:, :, y1:y2, x1:x2]\n        patch2 = F.pad(patch2, (x1, W - x2, y1, H - y2))  # outer-pad\n        patch2 = images_rolled - patch2  # inner-pad = img - outer-pad\n\n        images = patch1 + patch2  # cutmix img\n\n        lam = 1.0 - (x2 - x1) * (y2 - y1) / (W * H)\n        labels = lam * labels + (1.0 - lam) * labels_rolled  # cutmix label\n\n        images = images.view(*image_shape)\n        labels = labels.view(*label_shape)\n\n        return images, labels\n\n    return cutmix\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-03-11T16:35:20.680585Z","iopub.execute_input":"2024-03-11T16:35:20.680881Z","iopub.status.idle":"2024-03-11T16:35:20.695983Z","shell.execute_reply.started":"2024-03-11T16:35:20.680858Z","shell.execute_reply":"2024-03-11T16:35:20.695055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Pipeline\n* Reads the raw file and then decodes it to tf.Tensor\n* Resizes the image in desired size\n* Chages the datatype to **float32**\n* Caches the Data for boosting up the speed.\n* Uses Augmentations to reduce overfitting and make model more robust.\n* Finally, splits the data into batches.\n","metadata":{}},{"cell_type":"code","source":"# def build_decoder(with_labels=True, target_size=CFG.img_size, ext='png'):\n#     def decode_image(path):\n#         file_bytes = tf.io.read_file(path)\n#         if ext == 'png':\n#             img = tf.image.decode_png(file_bytes, channels=3, dtype=tf.uint8)\n#         elif ext in ['jpg', 'jpeg']:\n#             img = tf.image.decode_jpeg(file_bytes, channels=3)\n#         else:\n#             raise ValueError(\"Image extension not supported\")\n\n#         img = tf.image.resize(img, target_size, method='bilinear')\n#         img = tf.cast(img, tf.float32) / 255.0\n#         img = tf.reshape(img, [*target_size, 3])\n\n#         return img\n    \n#     def decode_label(label):\n#         label = tf.cast(label, tf.float32)\n#         return (label[0:1], label[1:2], label[2:5], label[5:8], label[8:11])\n    \n#     def decode_with_labels(path, label):\n#         return decode_image(path), decode_label(label)\n    \n#     return decode_with_labels if with_labels else decode\nimport cv2\nimport torch\nfrom torchvision import transforms\n\ndef build_decoder(with_labels=True, target_size=(CFG.img_size, CFG.img_size), ext='png'):\n    def decode_image(path):\n        img = cv2.imread(path, cv2.IMREAD_COLOR)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = cv2.resize(img, target_size, interpolation=cv2.INTER_LINEAR)\n        img = transforms.ToTensor()(img)\n        return img\n\n    def decode_label(label):\n        label = label.type(torch.float32)\n        return (label[0:1], label[1:2], label[2:5], label[5:8], label[8:11])\n\n    def decode_with_labels(path, label):\n        return decode_image(path), decode_label(label)\n\n    return decode_with_labels if with_labels else decode_image\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.697138Z","iopub.execute_input":"2024-03-11T16:35:20.698362Z","iopub.status.idle":"2024-03-11T16:35:20.87981Z","shell.execute_reply.started":"2024-03-11T16:35:20.698337Z","shell.execute_reply":"2024-03-11T16:35:20.879015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def build_augmenter(with_labels=True, dim=CFG.img_size):\n#     def augment(img, dim=dim):\n#         if random_float() < CFG.transform:\n#             img = transform(img,DIM=dim)\n#         img = tf.image.random_flip_left_right(img) if CFG.hflip else img\n#         img = tf.image.random_flip_up_down(img) if CFG.vflip else img\n#         if random_float() < CFG.pixel_aug:\n#             img = tf.image.random_hue(img, CFG.hue)\n#             img = tf.image.random_saturation(img, CFG.sat[0], CFG.sat[1])\n#             img = tf.image.random_contrast(img, CFG.cont[0], CFG.cont[1])\n#             img = tf.image.random_brightness(img, CFG.bri)\n#         img = tf.clip_by_value(img, 0, 1)  if CFG.clip else img         \n#         img = tf.reshape(img, [*dim, 3])\n#         return img\n    \n#     def augment_with_labels(img, label):    \n#         return augment(img), label\n    \n#     return augment_with_labels if with_labels else augment\n\nimport torchvision.transforms as transforms\nimport torch\nimport random\n\ndef build_augmenter(with_labels=True, dim=CFG.img_size):\n    def augment(img, dim=dim):\n        if random.random() < CFG.transform:\n            img = transform(img, DIM=dim)\n        img = transforms.functional.hflip(img) if CFG.hflip and random.random() < 0.5 else img\n        img = transforms.functional.vflip(img) if CFG.vflip and random.random() < 0.5 else img\n\n        if random.random() < CFG.pixel_aug:\n            img = transforms.ColorJitter(\n                brightness=CFG.bri,\n                contrast=(CFG.cont[0], CFG.cont[1]),\n                saturation=(CFG.sat[0], CFG.sat[1]),\n                hue=CFG.hue\n            )(img)\n\n        img = transforms.ToTensor()(img)\n        img = transforms.functional.normalize(img, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        return img\n\n    def augment_with_labels(img, label):\n        return augment(img), label\n\n    return augment_with_labels if with_labels else augment\n\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.881012Z","iopub.execute_input":"2024-03-11T16:35:20.881319Z","iopub.status.idle":"2024-03-11T16:35:20.891001Z","shell.execute_reply.started":"2024-03-11T16:35:20.881293Z","shell.execute_reply":"2024-03-11T16:35:20.890203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport os\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.8919Z","iopub.execute_input":"2024-03-11T16:35:20.892187Z","iopub.status.idle":"2024-03-11T16:35:20.903128Z","shell.execute_reply.started":"2024-03-11T16:35:20.892162Z","shell.execute_reply":"2024-03-11T16:35:20.902319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport cv2\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\n\nclass CustomDataset(Dataset):\n    def __init__(self, paths, labels=None, transform=None, augment_fn=None):\n        self.paths = paths\n        self.labels = labels\n        self.transform = transform\n        self.augment_fn = augment_fn\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, index):\n        path = self.paths[index]\n#         print('path:',path)\n        img = cv2.imread(path)  # Read image using cv2\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)  # Convert BGR to RGB\n#         img = cv2.resize(img,(256, 256))\n        if self.transform:\n            img = self.transform(img)\n\n        if self.augment_fn:\n            img = self.augment_fn(img)\n\n        if self.labels is not None:\n            label = self.labels[index]\n            return img, label\n        else:\n            return img\n\ndef build_transforms():\n    return transforms.Compose([\n        # Add your transformations here\n        transforms.ToTensor(),  # Convert numpy array to PyTorch tensor\n    ])\ndef build_augmenter():\n    def augment(img):\n        img = np.asarray(img)\n        img = cv2.flip(img, 1)\n        return img\n\n    return augment\n\n# Example Usage:\n# paths, labels = ... # Define your paths and labels\ntransform = build_transforms()\naugment_fn = build_augmenter()\n\n\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-03-11T16:35:20.904311Z","iopub.execute_input":"2024-03-11T16:35:20.904719Z","iopub.status.idle":"2024-03-11T16:35:20.917597Z","shell.execute_reply.started":"2024-03-11T16:35:20.904688Z","shell.execute_reply":"2024-03-11T16:35:20.916758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nfold = 0\nfold_df = df[df.fold==fold].sample(frac=1.0)\npaths  = fold_df.image_path.tolist()\nlabels = fold_df[CFG.target_col].values\n\n# Instantiate the CustomDataset\ndataset = CustomDataset(paths, labels, transform=transform, augment_fn=augment_fn)\n\n\n# Example DataLoader creation\nbatch_size = 1\ndataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)\n\n# Accessing a batch\ndata_iter = iter(dataloader)\nbatch = next(data_iter)","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:20.918664Z","iopub.execute_input":"2024-03-11T16:35:20.919006Z","iopub.status.idle":"2024-03-11T16:35:21.036028Z","shell.execute_reply.started":"2024-03-11T16:35:20.918975Z","shell.execute_reply":"2024-03-11T16:35:21.034988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# import matplotlib.pyplot as plt\n# import numpy as np\n\n# def display_batch(batch, size=2):\n#     imgs, labels = batch\n\n#     plt.figure(figsize=(size * 5, 10))\n\n#     for img_idx in range(size):\n#         plt.subplot(1, size, img_idx + 1)\n#         plt.title(f'Label: {labels[img_idx]}', fontsize=12)\n#         img = np.transpose(imgs[img_idx], (1, 2, 0))  # Convert (C, H, W) to (H, W, C) for display\n#         plt.imshow(img)\n#         plt.xticks([]); plt.yticks([])\n\n#     plt.tight_layout()\n#     plt.show()\n\n# # Example Usage:\n# display_batch(batch)\nimport matplotlib.pyplot as plt\nimport numpy as np\n\ndef display_batch(batch, size=2):\n    imgs, labels = batch\n\n    plt.figure(figsize=(size * 5, 10))\n\n    if not isinstance(labels, (list, np.ndarray)):\n        labels = [labels] * size\n\n    for img_idx in range(size):\n        plt.subplot(1, size, img_idx + 1)\n        plt.title(f'Label: {labels[img_idx]}', fontsize=12)\n        img = np.transpose(imgs[img_idx], (1, 2, 0))  # Convert (C, H, W) to (H, W, C) for display\n        plt.imshow(img)\n        plt.xticks([]); plt.yticks([])\n\n    plt.tight_layout()\n    plt.show()\n\n# Example Usage:\ndisplay_batch(batch)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T14:58:20.96343Z","iopub.execute_input":"2024-03-11T14:58:20.963901Z","iopub.status.idle":"2024-03-11T14:58:21.594173Z","shell.execute_reply.started":"2024-03-11T14:58:20.963846Z","shell.execute_reply":"2024-03-11T14:58:21.592797Z"}}},{"cell_type":"markdown","source":"## Dropout","metadata":{}},{"cell_type":"markdown","source":"import torch\nimport matplotlib.pyplot as plt\nimport numpy as np\n\ndef dropout(image, DIM, PROBABILITY=0.6, CT=5, SZ=0.1):\n    # input image - is one image of size [3, dim, dim] not a batch of [b, 3, dim, dim]\n    # output - image with CT squares of side size SZ*DIM removed\n\n    # DO DROPOUT WITH PROBABILITY DEFINED ABOVE\n    P = torch.randint(0, 2, (1,), dtype=torch.int32).item()\n    if (P == 0) or (CT == 0) or (SZ == 0):\n        return image\n\n    for k in range(CT):\n        # CHOOSE RANDOM LOCATION\n        x = torch.randint(0, DIM[1], (1,), dtype=torch.int32).item()\n        y = torch.randint(0, DIM[0], (1,), dtype=torch.int32).item()\n        # COMPUTE SQUARE\n        WIDTH = int(SZ * min(DIM))\n        ya = max(0, y - WIDTH // 2)\n        yb = min(DIM[0], y + WIDTH // 2)\n        xa = max(0, x - WIDTH // 2)\n        xb = min(DIM[1], x + WIDTH // 2)\n        # DROPOUT IMAGE\n        one = image[:, :, 0:ya]\n        two = torch.zeros((3, DIM[0], yb - ya), dtype=image.dtype)\n        three = image[:, :, yb:DIM[1]]\n        middle = torch.cat([one, two, three], dim=2)\n        image = torch.cat([image[:, :, 0:xa], middle, image[:, :, xb:DIM[1]]], dim=2)\n        image = image.view(3, *DIM)\n\n    return image\n\ndef visualize_dropout(original, dropout_result):\n    fig, axes = plt.subplots(1, 2, figsize=(10, 5))\n\n    axes[0].imshow(np.transpose(original.numpy(), (1, 2, 0)))\n    axes[0].set_title('Original Image')\n    axes[0].axis('off')\n\n    axes[1].imshow(np.transpose(dropout_result.numpy(), (1, 2, 0)))\n    axes[1].set_title('Image with Dropout')\n    axes[1].axis('off')\n\n    plt.show()\n\n# Example Usage:\nDIM = CFG.img_size  # Set DIM value based on your configuration\ndimgs = dropout(batch[0], DIM=DIM, PROBABILITY=1.0, CT=10, SZ=0.08)\nvisualize_dropout(batch[0][0], dimgs[0])\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-03-11T14:45:41.911333Z","iopub.execute_input":"2024-03-11T14:45:41.91202Z","iopub.status.idle":"2024-03-11T14:45:42.054265Z","shell.execute_reply.started":"2024-03-11T14:45:41.911988Z","shell.execute_reply":"2024-03-11T14:45:42.052664Z"}}},{"cell_type":"markdown","source":"## Affine Transform","metadata":{}},{"cell_type":"markdown","source":"timgs = tf.map_fn(lambda img: transform(img,DIM=CFG.img_size), batch[0])\nttars = batch[1]\ndisplay_batch((timgs, ttars), 3);","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-03-11T14:45:43.158182Z","iopub.execute_input":"2024-03-11T14:45:43.158593Z","iopub.status.idle":"2024-03-11T14:45:44.582635Z","shell.execute_reply.started":"2024-03-11T14:45:43.158562Z","shell.execute_reply":"2024-03-11T14:45:44.580902Z"}}},{"cell_type":"markdown","source":"# Loss Function\nWe'll use Binary Cross Entropy loss for `bowel` and `extravasation` at the same time we'll use Categorical Cross Entropy for `kidney`, `liver` and `spleen`.\n<!-- ## BCE\nLoss Function for this notebook is **BCE: Binary Crossentropy** or **Focal** as the task is **binary classification**\n\n$$\\textrm{BCE}  = -\\frac{1}{N} \\sum_{i=1}^{N} y_i \\cdot log(\\hat{y}_i) + (1 - y_i) \\cdot log(1 - \\hat{y}_i)$$\n\nwhere $\\hat{y}_i$ is the **predicted** value and $y_i$ is the **original** value for each instance $i$.\n\n## Focal\n$$\\textrm{FL} = -\\alpha_{t}(1 - p_{t})^{\\gamma}\\log{p_{t}}$$\n\n$\\gamma$ controls the shape of the curve. The higher the value of $\\gamma$, the lower the loss for well-classified examples, so we could turn the attention of the model more towards ‘hard-to-classify examples. FL gives high weights to the rare class and small weights to the dominating or common class. These weights are referred to as $\\alpha$.\n\n\n## Code\n* `tf.keras.losses.BinaryCrossentropy`\n* `tfa.losses.SigmoidFocalCrossEntropy` -->","metadata":{}},{"cell_type":"markdown","source":"# Build Model","metadata":{}},{"cell_type":"markdown","source":"from keras_cv_attention_models import efficientnet\n\ndef build_model(model_name=CFG.model_name,\n                loss_name=CFG.loss,\n                dim=CFG.img_size,\n                compile_model=True,\n                include_top=False):         \n    \n    # Define backbone\n    base = getattr(efficientnet, model_name)(input_shape=(*dim,3),\n                                    pretrained='imagenet',\n                                    num_classes=0) # get base model (efficientnet), use imgnet weights\n\n    inp = base.inputs\n    x = base.output\n    x = tf.keras.layers.GlobalAveragePooling2D()(x) # use GAP to get pooling result form conv outputs\n\n    # Define 'necks' for each head\n    x_bowel = tf.keras.layers.Dense(32, activation='silu')(x)\n    x_extra = tf.keras.layers.Dense(32, activation='silu')(x)\n    x_liver = tf.keras.layers.Dense(32, activation='silu')(x)\n    x_kidney = tf.keras.layers.Dense(32, activation='silu')(x)\n    x_spleen = tf.keras.layers.Dense(32, activation='silu')(x)\n\n    # Define heads\n    out_bowel = tf.keras.layers.Dense(1, name='bowel', activation='sigmoid')(x_bowel) # use sigmoid to convert predictions to [0-1]\n    out_extra = tf.keras.layers.Dense(1, name='extra', activation='sigmoid')(x_extra) # use sigmoid to convert predictions to [0-1]\n    out_liver = tf.keras.layers.Dense(3, name='liver', activation='softmax')(x_liver) # use softmax for the liver head\n    out_kidney = tf.keras.layers.Dense(3, name='kidney', activation='softmax')(x_kidney) # use softmax for the kidney head\n    out_spleen = tf.keras.layers.Dense(3, name='spleen', activation='softmax')(x_spleen) # use softmax for the spleen head\n\n    # Combine outputs\n#     out = tf.keras.layers.Concatenate()([out_bowel, out_extra, \n#                                          out_liver, out_kidney, out_spleen])\n    out = [out_bowel, out_extra, out_liver, out_kidney, out_spleen]\n\n    # Create model\n    model = tf.keras.Model(inputs=inp, outputs=out)\n\n    \n    if compile_model:\n        # optimizer\n        opt = tf.keras.optimizers.Adam(learning_rate=0.0001)\n        # loss\n        loss = {\n            'bowel':tf.keras.losses.BinaryCrossentropy(label_smoothing=0.05),\n            'extra':tf.keras.losses.BinaryCrossentropy(label_smoothing=0.05),\n            'liver':tf.keras.losses.CategoricalCrossentropy(label_smoothing=0.05),\n            'kidney':tf.keras.losses.CategoricalCrossentropy(label_smoothing=0.05),\n            'spleen':tf.keras.losses.CategoricalCrossentropy(label_smoothing=0.05),\n        }\n        # metric\n        metrics = {\n            'bowel':['accuracy'],\n            'extra':['accuracy'],\n            'liver':['accuracy'],\n            'kidney':['accuracy'],\n            'spleen':['accuracy'],\n        }\n        # compile\n        model.compile(optimizer=opt,\n                      loss=loss,\n                      metrics=metrics)\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-03-11T15:52:26.179355Z","iopub.execute_input":"2024-03-11T15:52:26.179939Z","iopub.status.idle":"2024-03-11T15:52:26.546342Z","shell.execute_reply.started":"2024-03-11T15:52:26.179906Z","shell.execute_reply":"2024-03-11T15:52:26.544857Z"}}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision import models\n\nclass MultiHeadModel(nn.Module):\n    def __init__(self, model_name='efficientnet_b0', num_classes=[1, 1, 3, 3, 3]):\n        super(MultiHeadModel, self).__init__()\n\n        # Load pre-trained backbone\n        if 'efficientnet' in model_name:\n            backbone = getattr(models, model_name)(pretrained=True)\n            backbone = nn.Sequential(*list(backbone.children())[:-1])\n        else:\n            raise NotImplementedError(\"Add support for other backbones\")\n\n        self.backbone = backbone\n        in_features = self.backbone(torch.randn(1, 3, 256, 256)).view(-1).size(0)\n\n        self.head_bowel = nn.Sequential(\n            nn.Linear(in_features, 32),\n            nn.SiLU(),\n            nn.Linear(32, num_classes[0]),\n            nn.Sigmoid()  # 保留 Sigmoid 作为二进制分类问题的输出激活函数\n        )\n\n        self.head_extra = nn.Sequential(\n            nn.Linear(in_features, 32),\n            nn.SiLU(),\n            nn.Linear(32, num_classes[1]),\n            nn.Sigmoid()  # 保留 Sigmoid 作为二进制分类问题的输出激活函数\n        )\n\n        self.head_liver = nn.Sequential(\n            nn.Linear(in_features, 32),\n            nn.SiLU(),\n            nn.Linear(32, num_classes[2])\n        )\n\n        self.head_kidney = nn.Sequential(\n            nn.Linear(in_features, 32),\n            nn.SiLU(),\n            nn.Linear(32, num_classes[3])\n        )\n\n        self.head_spleen = nn.Sequential(\n            nn.Linear(in_features, 32),\n            nn.SiLU(),\n            nn.Linear(32, num_classes[4])\n        )\n    def forward(self, x):\n        # Backbone\n        x = self.backbone(x)\n\n        # Global Average Pooling\n        x = F.adaptive_avg_pool2d(x, 1).view(x.size(0), -1)\n\n        # Heads\n        out_bowel = self.head_bowel(x)\n        out_extra = self.head_extra(x)\n        out_liver = self.head_liver(x)\n        out_kidney = self.head_kidney(x)\n        out_spleen = self.head_spleen(x)\n\n        return [out_bowel, out_extra, out_liver, out_kidney, out_spleen]\n\n# Create an instance of the PyTorch model\nmodel = MultiHeadModel()\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:21.037368Z","iopub.execute_input":"2024-03-11T16:35:21.037669Z","iopub.status.idle":"2024-03-11T16:35:21.807847Z","shell.execute_reply.started":"2024-03-11T16:35:21.037645Z","shell.execute_reply":"2024-03-11T16:35:21.806755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Assuming you already have your MultiHeadModel defined\nimport torch.optim as optim\n# Instantiate the model\nmodel = MultiHeadModel()\n\n# Define optimizer, loss functions, and metrics\noptimizer = optim.Adam(model.parameters(), lr=0.0001)\n\n# Loss functions\ncriterion_bowel = nn.BCELoss()\ncriterion_extra = nn.BCELoss()\ncriterion_liver = nn.CrossEntropyLoss()\ncriterion_kidney = nn.CrossEntropyLoss()\ncriterion_spleen = nn.CrossEntropyLoss()\n\n# Metrics\ndef accuracy(output, target):\n    _, predicted = torch.max(output, 1)\n    correct = (predicted == target).sum().item()\n    total = target.size(0)\n    acc = correct / total\n    return acc","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:21.809192Z","iopub.execute_input":"2024-03-11T16:35:21.809576Z","iopub.status.idle":"2024-03-11T16:35:22.056066Z","shell.execute_reply.started":"2024-03-11T16:35:21.809542Z","shell.execute_reply":"2024-03-11T16:35:22.055206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_lr_callback(batch_size=8, plot=False):\n    lr_start   = 0.000005\n    lr_max     = 0.00000050 * 10 * batch_size\n    lr_min     = 0.000001\n    lr_ramp_ep = 4\n    lr_sus_ep  = 0\n    lr_decay   = 0.8\n   \n    def lrfn(epoch):\n        if epoch < lr_ramp_ep:\n            lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start\n            \n        elif epoch < lr_ramp_ep + lr_sus_ep:\n            lr = lr_max\n            \n        elif CFG.scheduler=='exp':\n            lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min\n            \n        elif CFG.scheduler=='cosine':\n            decay_total_epochs = CFG.epochs - lr_ramp_ep - lr_sus_ep + 3\n            decay_epoch_index = epoch - lr_ramp_ep - lr_sus_ep\n            phase = math.pi * decay_epoch_index / decay_total_epochs\n            cosine_decay = 0.4 * (1 + math.cos(phase))\n            lr = (lr_max - lr_min) * cosine_decay + lr_min\n        return lr\n    if plot:\n        plt.figure(figsize=(10,5))\n        plt.plot(np.arange(CFG.epochs), [lrfn(epoch) for epoch in np.arange(CFG.epochs)], marker='o')\n        plt.xlabel('epoch'); plt.ylabel('learnig rate')\n        plt.title('Learning Rate Scheduler')\n        plt.show()\n\n    lr_callback = tf.keras.callbacks.LearningRateScheduler(lrfn, verbose=False)\n    return lr_callback\n\n_=get_lr_callback(CFG.batch_size, plot=True )","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-03-11T16:35:22.057441Z","iopub.execute_input":"2024-03-11T16:35:22.057882Z","iopub.status.idle":"2024-03-11T16:35:22.387263Z","shell.execute_reply.started":"2024-03-11T16:35:22.05785Z","shell.execute_reply":"2024-03-11T16:35:22.386345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# create directory to save gradcam imgs\n!mkdir -p gradcam\n\n# intialize wandb run\ndef wandb_init(fold):\n    config = {k:v for k,v in dict(vars(CFG)).items() if '__' not in k}\n    config.update({\"fold\":int(fold)})\n    yaml.dump(config, open(f'/kaggle/working/config fold-{fold}.yaml', 'w'),)\n    config = yaml.load(open(f'/kaggle/working/config fold-{fold}.yaml', 'r'), Loader=yaml.FullLoader)\n    run    = wandb.init(project=\"rsna-atd-public\",\n               name=f\"fold-{fold}|dim-{CFG.img_size[0]}x{CFG.img_size[1]}|model-{CFG.model_name}\",\n               config=config,\n               anonymous=anonymous,\n               group=CFG.exp_name\n                    )\n    return run\n\ndef log_wandb(fold):\n    \"log best result for error analysis\"\n    # log values to wandb\n    wandb.log({\n               'best_acc': best_acc,\n               'best_loss': best_loss,\n               'best_epoch': best_epoch,\n               'best_acc_bowel': best_acc_bowel,\n               'best_acc_extra': best_acc_extra,\n               'best_acc_liver': best_acc_liver,\n               'best_acc_kidney': best_acc_kidney,\n               'best_acc_spleen': best_acc_spleen,\n              })\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:22.388514Z","iopub.execute_input":"2024-03-11T16:35:22.388887Z","iopub.status.idle":"2024-03-11T16:35:23.373872Z","shell.execute_reply.started":"2024-03-11T16:35:22.388853Z","shell.execute_reply":"2024-03-11T16:35:23.372722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# get wandb callbacks\ndef get_wb_callbacks(fold):\n    wb_ckpt = wandb.keras.WandbModelCheckpoint(filepath='fold-%i.h5'%fold, \n                                               monitor='val_loss',\n                                               verbose=CFG.verbose,\n                                               save_best_only=True,\n                                               save_weights_only=False,\n                                               mode='min',)\n    wb_metr = wandb.keras.WandbMetricsLogger()\n    return [wb_ckpt, wb_metr]","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:23.375381Z","iopub.execute_input":"2024-03-11T16:35:23.375661Z","iopub.status.idle":"2024-03-11T16:35:23.381659Z","shell.execute_reply.started":"2024-03-11T16:35:23.375637Z","shell.execute_reply":"2024-03-11T16:35:23.380731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Model\n* Cross-Validation: 5 fold\n* **WandB** dashboard is shown end of the each fold. So we don't need to plot anything. We can select best model from here.","metadata":{}},{"cell_type":"markdown","source":"scores = []\n\nfor fold in np.arange(CFG.folds):\n    \n    # ignore not selected folds\n    if fold not in CFG.selected_folds:\n        continue\n        \n    # init wandb\n    if CFG.wandb:\n        run = wandb_init(fold)\n        wb_callbacks = get_wb_callbacks(fold)\n            \n    # train and valid dataframe\n    train_df = df.query(\"fold!=@fold\")\n    valid_df = df.query(\"fold==@fold\")\n    \n    # get image_paths and labels\n    train_paths = train_df.image_path.values; train_labels = train_df[CFG.target_col].values.astype(np.float32)\n    valid_paths = valid_df.image_path.values; valid_labels = valid_df[CFG.target_col].values.astype(np.float32)\n    test_paths  = test_df.image_path.values\n    \n    # shuffle train data\n    index = np.arange(len(train_df))\n    np.random.shuffle(index)\n    train_paths  = train_paths[index]\n    train_labels = train_labels[index]\n    \n    # min samples in debug mode\n    min_samples = CFG.batch_size*REPLICAS*2\n    \n    # for debug model run on small portion\n    if CFG.debug:\n        train_paths = train_paths[:min_samples]; train_labels = train_labels[:min_samples]\n        valid_paths = valid_paths[:min_samples]; valid_labels = valid_labels[:min_samples]\n    \n    # show message\n    print('#'*40); print('#### FOLD: ',fold)\n    print('#### IMAGE_SIZE: (%i, %i) | MODEL_NAME: %s | BATCH_SIZE: %i'%\n          (CFG.img_size[0],CFG.img_size[1],CFG.model_name,CFG.batch_size*REPLICAS))\n    \n    # data stat\n    num_train = len(train_paths)\n    num_valid = len(valid_paths)\n    if CFG.wandb:\n        wandb.log({'num_train':num_train,\n                   'num_valid':num_valid})\n    print('#### NUM_TRAIN: {:,} | NUM_VALID: {:,}'.format(num_train, num_valid))\n    \n    # build model\n    K.clear_session()\n    with strategy.scope():\n        model = build_model(CFG.model_name, dim=CFG.img_size, compile_model=True)\n\n    # build dataset\n    cache = 1 if 'TPU' in CFG.device else 0\n    train_ds = build_dataset(train_paths, train_labels, cache=cache, batch_size=CFG.batch_size*REPLICAS,\n                   repeat=True, shuffle=True, augment=CFG.augment)\n    val_ds = build_dataset(valid_paths, valid_labels, cache=cache, batch_size=CFG.batch_size*REPLICAS,\n                   repeat=False, shuffle=False, augment=False)\n    print('#'*40)   \n    \n    # callbacks\n    callbacks = []\n    ## save best model after each fold\n    sv = tf.keras.callbacks.ModelCheckpoint(\n        'fold-%i.h5'%fold, monitor='val_loss', verbose=CFG.verbose, save_best_only=True,\n        save_weights_only=False, mode='min', save_freq='epoch')\n    callbacks +=[sv]\n    ## lr-scheduler\n    callbacks += [get_lr_callback(CFG.batch_size)]\n    ## wandb callbacks\n    if CFG.wandb:\n        callbacks += wb_callbacks\n        \n    # train\n    print('Training...')\n    history = model.fit(\n        train_ds, \n        epochs=CFG.epochs if not CFG.debug else 2, \n        callbacks = callbacks, \n        steps_per_epoch=len(train_paths)/CFG.batch_size//REPLICAS,\n        validation_data=val_ds, \n        verbose=CFG.verbose\n    )\n    \n    # store best results\n    best_epoch = np.argmin(history.history['val_loss'])\n    best_loss = history.history['val_loss'][best_epoch]\n    best_acc_bowel = history.history['val_bowel_accuracy'][best_epoch]\n    best_acc_extra = history.history['val_extra_accuracy'][best_epoch]\n    best_acc_liver = history.history['val_liver_accuracy'][best_epoch]\n    best_acc_kidney = history.history['val_kidney_accuracy'][best_epoch]\n    best_acc_spleen = history.history['val_spleen_accuracy'][best_epoch]\n\n    # Find mean accuracy\n    best_acc = np.mean([best_acc_bowel, best_acc_extra, \n                        best_acc_liver, best_acc_kidney, best_acc_spleen])\n\n    print(f'\\n{\"=\"*17} FOLD {fold} RESULTS {\"=\"*17}')\n    print(f'>>>> BEST Loss  : {best_loss:.3f}\\n>>>> BEST Acc   : {best_acc:.3f}\\n>>>> BEST Epoch : {best_epoch}\\n')\n    print('ORGAN Acc:')\n    print(f'  >>>> {\"Bowel\".ljust(15)} : {best_acc_bowel:.3f}')\n    print(f'  >>>> {\"Extravasation\".ljust(15)} : {best_acc_extra:.3f}')\n    print(f'  >>>> {\"Liver\".ljust(15)} : {best_acc_liver:.3f}')\n    print(f'  >>>> {\"Kidney\".ljust(15)} : {best_acc_kidney:.3f}')\n    print(f'  >>>> {\"Spleen\".ljust(15)} : {best_acc_spleen:.3f}')\n\n    print(f'{\"=\"*50}\\n')\n\n    scores.append([best_loss, best_acc, \n                   best_acc_bowel, best_acc_extra, \n                   best_acc_liver, best_acc_kidney, best_acc_spleen])\n    \n    # log best result on wandb & plot\n    if CFG.wandb:\n        log_wandb(fold) # log\n        wandb.run.finish() # finish the run\n        display(ipd.IFrame(run.url, width=1080, height=720)) # show wandb dashboard","metadata":{"_kg_hide-input":true,"_kg_hide-output":false,"execution":{"iopub.status.busy":"2024-03-11T14:46:17.698978Z","iopub.execute_input":"2024-03-11T14:46:17.69941Z","iopub.status.idle":"2024-03-11T14:46:52.744988Z","shell.execute_reply.started":"2024-03-11T14:46:17.69936Z","shell.execute_reply":"2024-03-11T14:46:52.743178Z"}}},{"cell_type":"code","source":"################\n################ttttttooooorrrrrccccchh \n# h h\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom torchvision import transforms\nfrom tqdm import tqdm\n\n# Assuming you have the PyTorch equivalent of your model, dataset, and data loading functions\n\n# Initialize your model, loss function, optimizer, etc.\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = MultiHeadModel()  # Your PyTorch model\nmodel = model.to(device)\n\noptimizer = optim.Adam(model.parameters(), lr=0.0001)\n\n# Training loop\nscores = []\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:35:23.382883Z","iopub.execute_input":"2024-03-11T16:35:23.383251Z","iopub.status.idle":"2024-03-11T16:35:23.866828Z","shell.execute_reply.started":"2024-03-11T16:35:23.383215Z","shell.execute_reply":"2024-03-11T16:35:23.866007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nfor fold in range(CFG.folds):\n    if fold not in CFG.selected_folds:\n        continue\n\n    train_df = df.query(\"fold!=@fold\")\n    valid_df = df.query(\"fold==@fold\")\n\n    train_paths = train_df.image_path.values\n    train_labels = train_df[CFG.target_col].values.astype(np.float32)\n\n    valid_paths = valid_df.image_path.values\n    valid_labels = valid_df[CFG.target_col].values.astype(np.float32)\n\n    # Use PyTorch DataLoader for easier batching and shuffling\n    train_dataset = CustomDataset(train_paths, train_labels, transform=build_transforms(), augment_fn=build_augmenter())\n    train_loader = DataLoader(train_dataset, batch_size=CFG.batch_size, shuffle=True, num_workers=4)\n\n    valid_dataset = CustomDataset(valid_paths, valid_labels, transform=build_transforms())\n    valid_loader = DataLoader(valid_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=4)\n\n    # ... (rest of your setup code)\n\n    print('#' * 40)\n    print('#### FOLD: ', fold)\n    print('#### IMAGE_SIZE: (%i, %i) | MODEL_NAME: %s | BATCH_SIZE: %i' %\n          (CFG.img_size[0], CFG.img_size[1], CFG.model_name, CFG.batch_size))\n    print('#### NUM_TRAIN: {:,} | NUM_VALID: {:,}'.format(len(train_paths), len(valid_paths)))\n    print('#' * 40)\n\n    model.train()  # Set the model to training mode\n\n    for epoch in range(CFG.epochs if not CFG.debug else 2):\n        train_loss = 0.0\n\n        for inputs, labels in tqdm(train_loader, desc=f\"Epoch {epoch + 1}/{CFG.epochs}\"):\n            inputs, labels = inputs.to(device), labels.to(device)\n            optimizer.zero_grad()\n\n            outputs = model(inputs)\n            loss_bowel = nn.BCELoss()(outputs[0], labels[:,0].view(-1, 1))  # 二进制分类损失\n            loss_extra = nn.BCELoss()(outputs[1], labels[:,1].view(-1, 1))  # 二进制分类损失\n            loss_liver = nn.CrossEntropyLoss()(outputs[2], labels[:,2:5])  # 多分类损失，需要使用长整型标签\n            loss_kidney = nn.CrossEntropyLoss()(outputs[3], labels[:,5:8])  # 多分类损失，需要使用长整型标签\n            loss_spleen = nn.CrossEntropyLoss()(outputs[4], labels[:,8:11])  # 多分类损失，需要使用长整型标签\n\n\n            loss = loss_bowel + loss_extra + loss_liver + loss_kidney + loss_spleen\n\n            loss.backward()\n            optimizer.step()\n\n            train_loss += loss.item()\n# Print training loss at the end of each epoch\n        print(f\"Epoch {epoch + 1}/{CFG.epochs}, Training Loss: {train_loss / len(train_loader)}\")\n\n        # ... (rest of your evaluation and logging code)\n    # 保存模型参数和优化器状态\n    \n#         save_path = os.path.join('./train', f\"model_epoch_{epoch + 1}.pth\")\n        torch.save({\n            'epoch': epoch + 1,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'loss': train_loss / len(train_loader),\n        },f\"model_epoch_{epoch + 1}.pth\")\n# ... (rest of your code)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:46:29.272947Z","iopub.execute_input":"2024-03-11T16:46:29.273615Z","iopub.status.idle":"2024-03-11T16:49:10.505078Z","shell.execute_reply.started":"2024-03-11T16:46:29.273581Z","shell.execute_reply":"2024-03-11T16:49:10.503338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Calculate OOF Score","metadata":{}},{"cell_type":"code","source":"# overall oof pF1\noof_loss, oof_acc, oof_acc_bowel, oof_acc_extra, oof_acc_liver, oof_acc_kidney, oof_acc_spleen = np.array(scores).mean(axis=0)\n\nprint(f'\\n{\"=\"*15} OVERALL OOF RESULTS {\"=\"*15}')\nprint(f'>>>> OOF BEST Loss : {oof_loss:.3f}\\n>>>> OOF BEST Acc  : {oof_acc:.3f}\\n')\nprint('ORGAN OOF Acc:')\nprint(f'  >>>> {\"Bowel\".ljust(15)} : {oof_acc_bowel:.3f}')\nprint(f'  >>>> {\"Extravasation\".ljust(15)} : {oof_acc_extra:.3f}')\nprint(f'  >>>> {\"Liver\".ljust(15)} : {oof_acc_liver:.3f}')\nprint(f'  >>>> {\"Kidney\".ljust(15)} : {oof_acc_kidney:.3f}')\nprint(f'  >>>> {\"Spleen\".ljust(15)} : {oof_acc_spleen:.3f}')\nprint(f'{\"=\"*50}\\n')\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-03-11T16:38:43.741015Z","iopub.status.idle":"2024-03-11T16:38:43.741388Z","shell.execute_reply.started":"2024-03-11T16:38:43.741196Z","shell.execute_reply":"2024-03-11T16:38:43.74121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Remove Files","metadata":{}},{"cell_type":"code","source":"!rm -r /kaggle/working/wandb","metadata":{"execution":{"iopub.status.busy":"2024-03-11T16:38:43.742904Z","iopub.status.idle":"2024-03-11T16:38:43.743427Z","shell.execute_reply.started":"2024-03-11T16:38:43.743148Z","shell.execute_reply":"2024-03-11T16:38:43.743166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reference:\n1. [RANZCR: EfficientNet TPU Training](https://www.kaggle.com/xhlulu/ranzcr-efficientnet-tpu-training)\n1. [Triple Stratified KFold with TFRecords](https://www.kaggle.com/cdeotte/triple-stratified-kfold-with-tfrecords)","metadata":{}}]}