{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9114181,"sourceType":"datasetVersion","datasetId":5501168}],"dockerImageVersionId":30746,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport pydicom\nimport cv2\n\nimport math\nimport random\nimport numpy as np\nimport warnings\nimport os, gc, sys, copy, pickle\nimport glob\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\ntqdm.pandas()\n\nfrom joblib import Parallel, delayed\nimport multiprocessing as mp\n\nimport albumentations as A\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.cuda.amp as amp\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import WeightedRandomSampler\nfrom sklearn.utils.class_weight import compute_class_weight\n\nimport timm","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:01:23.059327Z","iopub.execute_input":"2024-08-06T02:01:23.059737Z","iopub.status.idle":"2024-08-06T02:01:39.116354Z","shell.execute_reply.started":"2024-08-06T02:01:23.059706Z","shell.execute_reply":"2024-08-06T02:01:39.115245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG = dict(\n    n_slice_per_c = 5, # number of slices per case\n    project_name = \"RSNA-2024-Lumbar-Spine-Classification-25D\",\n    artifact_name = \"rsnaEffNetModel\",\n    img_size = 256,\n    in_chans = 3, # color channels 3 for RGB\n    n_folds = 5,\n    backbone = \"efficientnet_b0.ra_in1k\",\n    p_rand_order_v1 = 0.2, # probability for applying randomization,\n    batch_size = 16,\n    out_dim = 3,\n    epochs = 50,\n    drop_rate = 0.,\n    drop_rate_last = 0.3,\n    drop_path_rate = 0.,\n    lr = 5e-5, # 1e-3, 8e-4, 5e-4, 4e-4\n    device = torch.device(\"cuda:0\") if torch.cuda.is_available() else \"cpu\",\n    seed = 2024,\n    log_wandb = False,\n    with_clip = False,\n\n)","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:01:39.121447Z","iopub.execute_input":"2024-08-06T02:01:39.121800Z","iopub.status.idle":"2024-08-06T02:01:39.188109Z","shell.execute_reply.started":"2024-08-06T02:01:39.121765Z","shell.execute_reply":"2024-08-06T02:01:39.186841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_PATH = Path(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\")\nos.listdir(DATA_PATH)","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:01:39.195174Z","iopub.execute_input":"2024-08-06T02:01:39.195858Z","iopub.status.idle":"2024-08-06T02:01:39.211795Z","shell.execute_reply.started":"2024-08-06T02:01:39.195823Z","shell.execute_reply":"2024-08-06T02:01:39.210694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(DATA_PATH/\"train.csv\")\ntrain_desc = pd.read_csv(DATA_PATH/\"train_series_descriptions.csv\")\ntest_desc = pd.read_csv(DATA_PATH/\"test_series_descriptions.csv\")\ntrain_coor = pd.read_csv(DATA_PATH/\"train_label_coordinates.csv\")\nsub = pd.read_csv(DATA_PATH/\"sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:01:39.213366Z","iopub.execute_input":"2024-08-06T02:01:39.213766Z","iopub.status.idle":"2024-08-06T02:01:39.415650Z","shell.execute_reply.started":"2024-08-06T02:01:39.213732Z","shell.execute_reply":"2024-08-06T02:01:39.414563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def reshape_row(row):\n    data = {'study_id': [], 'condition': [], 'level': [], 'severity': [], 'temp': []}\n    for column, value in row.items():\n        if column not in ['study_id']:\n            parts = column.split('_')\n            condition = ' '.join(word.capitalize() for word in parts[:-2])\n            level = parts[-2].capitalize() + '/'+ parts[-1].capitalize()\n            data['study_id'].append(row['study_id'])\n            data['condition'].append(condition)\n            data['level'].append(level)\n            data['severity'].append(value)\n            data['temp'].append(column)\n\n    return pd.DataFrame(data)\n\n\nnew_train = pd.concat([reshape_row(row) for _, row in train.iterrows()], ignore_index=True)\nnew_train.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:01:39.419238Z","iopub.execute_input":"2024-08-06T02:01:39.420042Z","iopub.status.idle":"2024-08-06T02:01:41.119533Z","shell.execute_reply.started":"2024-08-06T02:01:39.419999Z","shell.execute_reply":"2024-08-06T02:01:41.118306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"merged_train = pd.merge(new_train, train_coor, on=['study_id', 'condition', 'level'], how='inner')\nfinal_merged_train = pd.merge(merged_train, train_desc, on=['study_id','series_id'], how='inner')\nfinal_merged_train['row_id'] = (final_merged_train['study_id']).astype(str) + '_' + final_merged_train['temp']\nfinal_merged_train['image_path'] = (\n    f'{str(DATA_PATH)}/train_images/' + \n    final_merged_train['study_id'].astype(str) + '/' +\n    final_merged_train['series_id'].astype(str) + '/' +\n    final_merged_train['instance_number'].astype(str) + '.dcm'\n)\nfinal_merged_train['severity'] = final_merged_train['severity'].map(\n    {'Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'}\n)\nfinal_merged_train.drop(['temp'], axis=1, inplace=True)\ntrain_data = final_merged_train.copy()\ntrain_data.head()","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:01:41.120874Z","iopub.execute_input":"2024-08-06T02:01:41.121195Z","iopub.status.idle":"2024-08-06T02:01:41.446742Z","shell.execute_reply.started":"2024-08-06T02:01:41.121167Z","shell.execute_reply":"2024-08-06T02:01:41.445476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def check_exists(path):\n    return os.path.exists(path)\n\ndef check_study_id(row):\n    study_id = row['study_id']\n    path = f'{str(DATA_PATH)}/train_images/{study_id}'\n    return check_exists(path)\n\ndef check_series_id(row):\n    study_id = row['study_id']\n    series_id = row['series_id']\n    path = f'{str(DATA_PATH)}/train_images/{study_id}/{series_id}'\n    return check_exists(path)\n\n# Define a function to check if an image file exists\ndef check_image_exists(row):\n    image_path = row['image_path']\n    return check_exists(image_path)\n\n\ntrain_data['study_id_exists'] = train_data.progress_apply(check_study_id, axis=1)\ntrain_data['series_id_exists'] = train_data.progress_apply(check_series_id, axis=1)\ntrain_data['image_exists'] = train_data.progress_apply(check_image_exists, axis=1)\n\ntrain_data = train_data[(train_data['study_id_exists']) & (train_data['series_id_exists']) & (train_data['image_exists'])]\ntrain_data.shape[0]","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:01:41.448129Z","iopub.execute_input":"2024-08-06T02:01:41.448558Z","iopub.status.idle":"2024-08-06T02:02:47.289620Z","shell.execute_reply.started":"2024-08-06T02:01:41.448529Z","shell.execute_reply":"2024-08-06T02:02:47.288329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label2id = {v: i for i,v in enumerate(train_data['severity'].unique())}\ntrain_data['target'] = train_data['severity'].map(label2id)\ntrain_data = train_data.dropna(subset=['severity']).reset_index(drop=True)\n\nseries2id = {v: i for i,v in enumerate(train_data['series_description'].unique())}\ntrain_data['series2id'] = train_desc['series_description'].map(series2id)\ntrain_data = train_data.dropna(subset=['series2id']).reset_index(drop=True)\ntrain_data[train_data['target'] == 2].head()","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.291110Z","iopub.execute_input":"2024-08-06T02:02:47.291520Z","iopub.status.idle":"2024-08-06T02:02:47.370997Z","shell.execute_reply.started":"2024-08-06T02:02:47.291484Z","shell.execute_reply":"2024-08-06T02:02:47.369906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seeding(SEED):\n    np.random.seed(SEED)\n    random.seed(SEED)\n    os.environ['PYTHONHASHSEED'] = str(SEED)\n    torch.manual_seed(SEED)\n    if torch.cuda.is_available(): \n        torch.cuda.manual_seed(SEED)\n        torch.cuda.manual_seed_all(SEED)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = True\n    print('seeding done!!!')\n\ndef flush():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.reset_peak_memory_stats()","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.376389Z","iopub.execute_input":"2024-08-06T02:02:47.376909Z","iopub.status.idle":"2024-08-06T02:02:47.384632Z","shell.execute_reply.started":"2024-08-06T02:02:47.376876Z","shell.execute_reply":"2024-08-06T02:02:47.383443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.385846Z","iopub.execute_input":"2024-08-06T02:02:47.386195Z","iopub.status.idle":"2024-08-06T02:02:47.397723Z","shell.execute_reply.started":"2024-08-06T02:02:47.386156Z","shell.execute_reply":"2024-08-06T02:02:47.396686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()\n\nclass LumbarDataset(Dataset):\n    def __init__(self, df, mode, transform, label_name='target'):\n        self.df = df\n        self.mode = mode\n        self.transform = transform\n        self.label = df.loc[:, label_name]\n\n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        study_id = self.df.loc[idx, 'study_id']\n        series_id = self.df.loc[idx, 'series_id']\n        target = self.df.loc[idx, 'target']\n        image_paths = glob.glob(f\"{str(DATA_PATH)}/train_images/{study_id}/{series_id}/*.dcm\")\n\n        if len(image_paths) > CONFIG['n_slice_per_c']:\n            random.shuffle(image_paths)\n            image_paths = sorted(image_paths)[:CONFIG['n_slice_per_c']]\n\n        # numpy 4D array to store preprocessed image\n        images = np.zeros((CONFIG['n_slice_per_c'], CONFIG['in_chans'], CONFIG['img_size'], CONFIG['img_size'])).astype(np.float32)\n        for idx in list(range(min(CONFIG['n_slice_per_c'], len(image_paths)))):\n            filepath = image_paths[idx]\n#             image = cv2.imread(filepath, 0) # grayscale image\n            image = load_dicom(filepath)\n            if CONFIG['in_chans'] == 1:\n                image = np.expand_dims(image, axis=-1) # adds a dimension to the grayscale image\n            elif CONFIG['in_chans'] == 3:\n                image = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR) # converts grayscale to BGR format\n\n            image = self.transform(image=image)['image']\n            image = image.transpose(2, 0, 1).astype(np.float32) / 255\n            images[idx, ...] = image\n\n        images = torch.tensor(images).float()\n        if self.mode != 'test':\n            labels = torch.tensor([target] * CONFIG['n_slice_per_c']).float()\n\n            if self.mode == 'train' and random.random() < CONFIG['p_rand_order_v1']:\n                indices = torch.randperm(images.size(0))\n                images = images[indices]\n            return images, labels\n        else:\n            return images\n        \n    def get_labels(self):\n        return self.label","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.399445Z","iopub.execute_input":"2024-08-06T02:02:47.400030Z","iopub.status.idle":"2024-08-06T02:02:47.610847Z","shell.execute_reply.started":"2024-08-06T02:02:47.399967Z","shell.execute_reply":"2024-08-06T02:02:47.609836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(height, width):\n    train_tsfm = A.Compose([\n        A.Resize(height=512, width=512),\n        # Geometric augementations\n        A.Perspective(p=0.5),\n        A.HorizontalFlip(p=0.3),\n        A.VerticalFlip(p=0.2),\n        A.SafeRotate((-25,25), p=0.5),\n        # A.Rotate(-25, 25, p=0.5),\n        # Color Augmentations\n        A.RandomBrightnessContrast(p=0.5),\n        #A.RandomContrast(limit=1.5, p=0.5),\n        #A.ColorJitter(saturation=0.1, hue=0.1, p=0.5),\n        #A.RandomGamma(gamma_limit=(80, 120), p=0.5),\n        A.CenterCrop(height=height, width=width, p=1.0)\n    ]) # training transformations\n\n    valid_tsfm = A.Compose([\n        A.Resize(height=height, width=width)\n    ]) # validation/evaluation transformation\n\n    return {\"train\": train_tsfm, \"eval\": valid_tsfm}\n\ndef get_dataloaders(data, cfg, split=\"train\"):\n    img_size = cfg['img_size']\n    height, width = img_size, img_size\n    tsfm = get_transforms(height, width)\n    if split == 'train':\n        train_tsfm = tsfm['train']\n        ds = LumbarDataset(data, mode='train', transform=train_tsfm)\n        labels = ds.get_labels()\n        class_weights = torch.tensor([1,2,4])\n        samples_weights = class_weights[labels]\n\n        sampler = WeightedRandomSampler(weights=samples_weights, num_samples=len(samples_weights), replacement=True)\n        dls = DataLoader(ds, batch_size=cfg['batch_size'],sampler=sampler,num_workers=os.cpu_count(),drop_last=True, pin_memory=True)\n\n    elif split == 'valid' or split == 'test':\n        eval_tsfm = tsfm['eval']\n        ds = LumbarDataset(data, mode='valid', transform=eval_tsfm)\n        dls = DataLoader(ds,batch_size=2*cfg['batch_size'], \n                         shuffle=False, \n                         num_workers=os.cpu_count(), \n                         drop_last=False, \n                         pin_memory=True)\n    else:\n        raise Exception(\"Split should be 'train' or 'valid' or 'test'\")\n    \n    return dls","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.612420Z","iopub.execute_input":"2024-08-06T02:02:47.612815Z","iopub.status.idle":"2024-08-06T02:02:47.631602Z","shell.execute_reply.started":"2024-08-06T02:02:47.612765Z","shell.execute_reply":"2024-08-06T02:02:47.630678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn import model_selection\n\nkfold = model_selection.StratifiedKFold(n_splits=5, shuffle=True, random_state=2024)\nx = train_data.index.values\ny = train_data['target'].values.astype(int)\n\ntrain_data['fold'] = -1\nfor fold, (tr_idx, val_idx) in enumerate(kfold.split(x,y)):\n    train_data.loc[val_idx, 'fold'] = fold\n\ntrain_data.groupby('fold')['target'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.633018Z","iopub.execute_input":"2024-08-06T02:02:47.633419Z","iopub.status.idle":"2024-08-06T02:02:47.673694Z","shell.execute_reply.started":"2024-08-06T02:02:47.633373Z","shell.execute_reply":"2024-08-06T02:02:47.672625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TimmModel(nn.Module):\n    def __init__(self, backbone, pretrained=False):\n        super(TimmModel, self).__init__()\n\n        self.encoder = timm.create_model(\n            backbone,\n            in_chans = CONFIG['in_chans'],\n            num_classes = CONFIG['out_dim'],\n            features_only = False,\n            drop_rate = CONFIG['drop_rate'],\n            drop_path_rate = CONFIG['drop_path_rate'],\n            pretrained=pretrained\n        )\n\n        if 'efficient' in backbone:\n            hdim = self.encoder.conv_head.out_channels\n            self.encoder.classifier = nn.Identity()\n        elif 'convnext' in backbone:\n            hdim = self.encoder.head.fc.in_features\n            self.encoder.head.fc = nn.Identity()\n\n\n        self.lstm = nn.LSTM(hdim, 256, num_layers=2, dropout=CONFIG[\"drop_rate\"], bidirectional=True, batch_first=True)\n        self.head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.Dropout(CONFIG[\"drop_rate_last\"]),\n            nn.LeakyReLU(0.1),\n            nn.Linear(256, CONFIG[\"out_dim\"]),\n        )\n\n    def forward(self, x):  # (bs, nslice, ch, sz, sz)\n        bs = x.shape[0]\n        x = x.view(bs * CONFIG[\"n_slice_per_c\"], CONFIG[\"in_chans\"], CONFIG[\"img_size\"], CONFIG[\"img_size\"])\n        feat = self.encoder(x)\n        feat = feat.view(bs, CONFIG[\"n_slice_per_c\"], -1)\n        feat, _ = self.lstm(feat)\n        feat = feat.contiguous().view(bs * CONFIG[\"n_slice_per_c\"], -1)\n        feat = self.head(feat)\n        feat = feat.view(bs, CONFIG[\"n_slice_per_c\"], CONFIG[\"out_dim\"]).contiguous()\n\n        return feat","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.674961Z","iopub.execute_input":"2024-08-06T02:02:47.675294Z","iopub.status.idle":"2024-08-06T02:02:47.689331Z","shell.execute_reply.started":"2024-08-06T02:02:47.675266Z","shell.execute_reply":"2024-08-06T02:02:47.688233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\n\ndef flush():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.reset_peak_memory_stats()\n\ndef init_weights(m):\n    if type(m) == nn.Linear:\n        torch.nn.init.xavier_uniform(m.weight)\n        m.bias.data.fill_(0.01)\n\ndef shared_step(model, batch, criterion):\n    image, target = batch[0], batch[1]\n    image = image.to(CONFIG[\"device\"], non_blocking=True)\n    target = target.to(CONFIG[\"device\"], non_blocking=True)\n    logits = model(image.to(torch.float32))\n    logits = logits.mean(axis=1)\n    target = target.mean(axis=-1).to(torch.int64)\n    loss = criterion(logits, target)\n#     loss = criterion(logits.view(-1, CONFIG[\"out_dim\"]), target.view(-1).to(torch.int64))\n\n    return {\n        \"loss\": loss,\n        \"logits\": logits,\n        \"target\": target,\n    }","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.690957Z","iopub.execute_input":"2024-08-06T02:02:47.691304Z","iopub.status.idle":"2024-08-06T02:02:47.703391Z","shell.execute_reply.started":"2024-08-06T02:02:47.691277Z","shell.execute_reply":"2024-08-06T02:02:47.702467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name = \"/kaggle/input/model-25/rsna_2024_lumbar_spine_fold_0_epoch_4.pth\"\ndef get_model():\n    model = TimmModel(backbone=CONFIG[\"backbone\"], pretrained=False)\n    model.load_state_dict(torch.load(model_name))\n    return model\n\n\ndef predict_test_data(testloader, expanded_test_desc):\n    predictions = []\n    normal_mild_probs = []\n    moderate_probs = []\n    severe_probs = []\n    \n    device = CONFIG['device']\n    model = get_model().to(device)\n    model.eval()  # Set the model to eval mode\n        \n    with torch.no_grad():\n        for idx, images in enumerate(tqdm(testloader)):\n            images = images.to(device)\n            outputs = model(images)\n            probs = torch.softmax(outputs, dim=2).mean(dim=1) \n            \n            for prob in probs:\n                normal_mild_probs.append(prob[0].item())\n                moderate_probs.append(prob[1].item())\n                severe_probs.append(prob[2].item())\n                predictions.append(prob)\n    flush()\n    return normal_mild_probs, moderate_probs, severe_probs, predictions\n","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.704609Z","iopub.execute_input":"2024-08-06T02:02:47.704899Z","iopub.status.idle":"2024-08-06T02:02:47.721721Z","shell.execute_reply.started":"2024-08-06T02:02:47.704874Z","shell.execute_reply":"2024-08-06T02:02:47.720895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images'\n\ndef get_image_paths(row): \n    series_path = os.path.join(base_path, str(row['study_id']), str(row['series_id']))\n    if os.path.exists(series_path):\n        return [os.path.join(series_path, f) for f in os.listdir(series_path) if os.path.isfile(os.path.join(series_path, f))]\n    return []\n\ncondition_mapping = {\n    'Sagittal T1': {'left': 'left_neural_foraminal_narrowing', 'right': 'right_neural_foraminal_narrowing'},\n    'Axial T2': {'left': 'left_subarticular_stenosis', 'right': 'right_subarticular_stenosis'},\n    'Sagittal T2/STIR': 'spinal_canal_stenosis'\n}\n\nexpanded_rows = []\n\nfor index, row in test_desc.iterrows():\n    image_paths = get_image_paths(row)\n    conditions = condition_mapping.get(row['series_description'], {})\n    if isinstance(conditions, str):  # Single condition\n        conditions = {'left': conditions, 'right': conditions}\n    for side, condition in conditions.items():\n        for image_path in image_paths:\n            expanded_rows.append({\n                'study_id': row['study_id'],\n                'series_id': row['series_id'],\n                'series_description': row['series_description'],\n                'image_path': image_path,\n                'condition': condition,\n                'row_id': f\"{row['study_id']}_{condition}\"\n            })\n\n# Create a new dataframe from the expanded rows\nexpanded_test_desc = pd.DataFrame(expanded_rows)\n\n# Display the resulting dataframe\nexpanded_test_desc.head(5)\n    ","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.722923Z","iopub.execute_input":"2024-08-06T02:02:47.723235Z","iopub.status.idle":"2024-08-06T02:02:47.778238Z","shell.execute_reply.started":"2024-08-06T02:02:47.723208Z","shell.execute_reply":"2024-08-06T02:02:47.777349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"levels = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n\n# Function to update row_id with levels\ndef update_row_id(row, levels):\n    level = levels[row.name % len(levels)]\n    return f\"{row['study_id']}_{row['condition']}_{level}\"\n\n# Update row_id in expanded_test_desc to include levels\nexpanded_test_desc['row_id'] = expanded_test_desc.apply(lambda row: update_row_id(row, levels), axis=1)\nexpanded_test_desc.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.779557Z","iopub.execute_input":"2024-08-06T02:02:47.779901Z","iopub.status.idle":"2024-08-06T02:02:47.799123Z","shell.execute_reply.started":"2024-08-06T02:02:47.779869Z","shell.execute_reply":"2024-08-06T02:02:47.798036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        image_path = self.dataframe['image_path'][index]\n        image = load_dicom(image_path)\n        \n        # Ensure image is 3 channel\n        if len(image.shape) == 2:  # If grayscale\n            image = np.stack([image] * 3, axis=-1)  # Convert to 3 channel\n        \n        if self.transform:\n            image = self.transform(image=image)['image']\n        \n        # Convert to tensor\n        image = torch.tensor(image).float()\n        \n        # Ensure CHW format\n        image = image.permute(2, 0, 1)\n        \n        # Normalize to [0, 1] if not already\n        if image.max() > 1:\n            image = image / 255.0\n        \n        # Repeat the slice to match n_slice_per_c\n        image = image.unsqueeze(0).repeat(CONFIG['n_slice_per_c'], 1, 1, 1)\n        \n        return image\n\n# Define the transforms\ntransform = A.Compose([\n    A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n])\n\n# Create a test dataset and dataloader\ntest_dataset = TestDataset(expanded_test_desc, transform)\ntestloader = DataLoader(test_dataset, batch_size=CONFIG['batch_size'], shuffle=False)\nflush()","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:47.800970Z","iopub.execute_input":"2024-08-06T02:02:47.801417Z","iopub.status.idle":"2024-08-06T02:02:48.040285Z","shell.execute_reply.started":"2024-08-06T02:02:47.801359Z","shell.execute_reply":"2024-08-06T02:02:48.039455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"normal_mild_probs, moderate_probs, severe_probs, test_predictions = predict_test_data(testloader, expanded_test_desc)","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:48.041362Z","iopub.execute_input":"2024-08-06T02:02:48.041680Z","iopub.status.idle":"2024-08-06T02:02:58.834732Z","shell.execute_reply.started":"2024-08-06T02:02:48.041654Z","shell.execute_reply":"2024-08-06T02:02:58.833748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_predictions[0]","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:58.836040Z","iopub.execute_input":"2024-08-06T02:02:58.836322Z","iopub.status.idle":"2024-08-06T02:02:59.253555Z","shell.execute_reply.started":"2024-08-06T02:02:58.836299Z","shell.execute_reply":"2024-08-06T02:02:59.252560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"expanded_test_desc['normal_mild'] = normal_mild_probs\nexpanded_test_desc['moderate'] = moderate_probs\nexpanded_test_desc['severe'] = severe_probs","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:59.254677Z","iopub.execute_input":"2024-08-06T02:02:59.254964Z","iopub.status.idle":"2024-08-06T02:02:59.261965Z","shell.execute_reply.started":"2024-08-06T02:02:59.254939Z","shell.execute_reply":"2024-08-06T02:02:59.260848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = expanded_test_desc[[\"row_id\",\"normal_mild\",\"moderate\",\"severe\"]]","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:59.263228Z","iopub.execute_input":"2024-08-06T02:02:59.263586Z","iopub.status.idle":"2024-08-06T02:02:59.275680Z","shell.execute_reply.started":"2024-08-06T02:02:59.263559Z","shell.execute_reply":"2024-08-06T02:02:59.274617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.head(10)","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:59.276995Z","iopub.execute_input":"2024-08-06T02:02:59.277702Z","iopub.status.idle":"2024-08-06T02:02:59.298479Z","shell.execute_reply.started":"2024-08-06T02:02:59.277670Z","shell.execute_reply":"2024-08-06T02:02:59.297391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"grouped_submission = submission.groupby('row_id').max().reset_index()\ngrouped_submission","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:59.299989Z","iopub.execute_input":"2024-08-06T02:02:59.300360Z","iopub.status.idle":"2024-08-06T02:02:59.330419Z","shell.execute_reply.started":"2024-08-06T02:02:59.300329Z","shell.execute_reply":"2024-08-06T02:02:59.329365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub[['normal_mild', 'moderate', 'severe']] = grouped_submission[['normal_mild', 'moderate', 'severe']]","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:59.331817Z","iopub.execute_input":"2024-08-06T02:02:59.332213Z","iopub.status.idle":"2024-08-06T02:02:59.339821Z","shell.execute_reply.started":"2024-08-06T02:02:59.332174Z","shell.execute_reply":"2024-08-06T02:02:59.338753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.to_csv(\"/kaggle/working/submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-08-06T02:02:59.344768Z","iopub.execute_input":"2024-08-06T02:02:59.345071Z","iopub.status.idle":"2024-08-06T02:02:59.357112Z","shell.execute_reply.started":"2024-08-06T02:02:59.345044Z","shell.execute_reply":"2024-08-06T02:02:59.356202Z"},"trusted":true},"execution_count":null,"outputs":[]}]}