{"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":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport matplotlib.pyplot as plt\nimport random\nfrom PIL import Image\nimport pydicom\nfrom sklearn.model_selection import train_test_split\nimport nibabel as nib\nimport seaborn as sns\nfrom tqdm import tqdm\n\nimport pytorch_lightning as L\nfrom pytorch_lightning.loggers import CSVLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\nfrom pytorch_lightning.loggers import TensorBoardLogger\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchmetrics\nfrom torch.utils.data import DataLoader\nfrom torchvision import datasets, transforms\nfrom torch.utils.data import Dataset\nimport torch.optim as optim","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-11-03T17:02:56.667296Z","iopub.execute_input":"2023-11-03T17:02:56.667938Z","iopub.status.idle":"2023-11-03T17:03:11.64025Z","shell.execute_reply.started":"2023-11-03T17:02:56.667903Z","shell.execute_reply":"2023-11-03T17:03:11.639456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    SEED = 101\n    BATCH_SIZE = 16\n    MAX_EPOCHS = 5\n    LR=0.005\n\nconfig = Config()\nrandom.seed(config.SEED)\nnp.random.seed(config.SEED)","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:03:11.642064Z","iopub.execute_input":"2023-11-03T17:03:11.642381Z","iopub.status.idle":"2023-11-03T17:03:11.647397Z","shell.execute_reply.started":"2023-11-03T17:03:11.642356Z","shell.execute_reply":"2023-11-03T17:03:11.646436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"original = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/train.csv')\ntest_or = pd.read_parquet('/kaggle/input/rsna-2023-abdominal-trauma-detection/test_dicom_tags.parquet')","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:03:11.648411Z","iopub.execute_input":"2023-11-03T17:03:11.648685Z","iopub.status.idle":"2023-11-03T17:03:11.892911Z","shell.execute_reply.started":"2023-11-03T17:03:11.648662Z","shell.execute_reply":"2023-11-03T17:03:11.892109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_injuried = original[original['any_injury']==1]\ninjuried = df_injuried['patient_id']","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:03:11.894969Z","iopub.execute_input":"2023-11-03T17:03:11.895304Z","iopub.status.idle":"2023-11-03T17:03:11.901728Z","shell.execute_reply.started":"2023-11-03T17:03:11.895279Z","shell.execute_reply":"2023-11-03T17:03:11.900923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Creation of the path of the DICOM images\nroot_dir = '/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images'\npatients = injuried\n\nserial_path = {'patient_id':[],\n               'serial_id':[],\n               'serial_path':[]\n              }\n\nserials = {'patient_path':[],\n           'patient': []\n          }\n# Populate serials['patient_path'] first\nfor patient_id in patients:\n    patient_path = os.path.join(root_dir, str(patient_id))\n    serials['patient_path'].append(patient_path)\n    serials['patient'].append(patient_id)\n    \nserial_list = list(serials.values())\npatient_paths, patient_ids = serial_list\n\n#Now, let's create entries in serial_path for each serial ID\nfor patient_path, patient_id in zip(patient_paths, patient_ids):\n    serial_ids = os.listdir(patient_path)\n    for serial_id in serial_ids:\n        serial_path['patient_id'].append(patient_id)\n        serial_path['serial_id'].append(serial_id)\n        serial_path['serial_path'].append(os.path.join(patient_path,serial_id))\n\ndf_serials = pd.DataFrame(serial_path)\n\nimage_paths = {'patient_id':[],'serial_id':[],'serial_path':[],'image_path':[], 'image':[] }\n\npatient_path_dir = df_serials['patient_id']\nserial_id_dir = df_serials['serial_id']\nserial_path_dir = df_serials['serial_path']\n\nfor patient_id, serial_id, serial_path in zip(patient_path_dir, serial_id_dir, serial_path_dir):\n    images = os.listdir(serial_path)\n    for image in images:\n        image_paths['patient_id'].append(patient_id)\n        image_paths['serial_id'].append(serial_id)\n        image_paths['image_path'].append(os.path.join(serial_path,image))\n        image_paths['serial_path'].append(serial_path)\n        image_paths['image'].append(image)\n        \nimage_paths_df = pd.DataFrame(image_paths)\n\nimage_paths_df['image'] = image_paths_df['image'].str.replace('.dcm', '').astype(int)","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:03:11.903043Z","iopub.execute_input":"2023-11-03T17:03:11.903376Z","iopub.status.idle":"2023-11-03T17:06:21.897828Z","shell.execute_reply.started":"2023-11-03T17:03:11.903346Z","shell.execute_reply":"2023-11-03T17:06:21.896898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Separate the serial_id with more than 50 images, less than 50 images and exactly 50 images\npidf = pd.DataFrame()\nlidf = pd.DataFrame()\nmidf = pd.DataFrame()\nnum_slices = 50\nfor serial, group in image_paths_df.groupby('serial_id'):\n    if len(group) == num_slices:\n        group_sorted = group.sort_values('image')\n        pidf = pd.concat([pidf,group_sorted],axis=0,ignore_index=True)\n    elif len(group) > num_slices:\n        group_sorted = group.sort_values('image')\n        midf = pd.concat([midf,group_sorted],axis=0,ignore_index=True)\n    else:\n        group_sorted = group.sort_values('image')\n        lidf = pd.concat([lidf,group_sorted],axis=0,ignore_index=True)","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:06:21.899168Z","iopub.execute_input":"2023-11-03T17:06:21.899802Z","iopub.status.idle":"2023-11-03T17:06:34.212911Z","shell.execute_reply.started":"2023-11-03T17:06:21.899767Z","shell.execute_reply":"2023-11-03T17:06:34.2119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_more_df = pd.DataFrame()  \nfor serial_id, group in midf.groupby('serial_id'):\n    total_rows = group.shape[0]\n    step_size = max(total_rows // num_slices, 1)  \n    selected_rows = group.iloc[::step_size][:num_slices].reset_index(drop=True)\n    new_more_df = pd.concat([new_more_df, selected_rows], ignore_index=True)","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:06:34.214171Z","iopub.execute_input":"2023-11-03T17:06:34.214465Z","iopub.status.idle":"2023-11-03T17:06:36.287184Z","shell.execute_reply.started":"2023-11-03T17:06:34.21444Z","shell.execute_reply":"2023-11-03T17:06:36.286179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_less_df = pd.DataFrame()\nfor serial, group in lidf.groupby('serial_id'):\n    length = len(group)\n    intergers = num_slices//length\n    rest = num_slices%length\n    group_df = pd.concat([group] * intergers, axis=0)\n    if rest >0:\n        step_size = max(length // rest, 1)\n        selected_rows = group.iloc[::step_size][:rest]\n        inter_df = pd.concat([group_df,selected_rows], axis=0)\n    new_less_df = pd.concat([new_less_df,inter_df], axis=0)\n    new_less_df = new_less_df.sort_index()","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:06:36.288396Z","iopub.execute_input":"2023-11-03T17:06:36.288696Z","iopub.status.idle":"2023-11-03T17:06:36.308775Z","shell.execute_reply.started":"2023-11-03T17:06:36.288672Z","shell.execute_reply":"2023-11-03T17:06:36.307873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_dfs = pd.concat([pidf,new_less_df,new_more_df],axis=0,ignore_index=True)\ndf_total = pd.merge(total_dfs,df_injuried, on='patient_id')","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:06:36.310217Z","iopub.execute_input":"2023-11-03T17:06:36.310891Z","iopub.status.idle":"2023-11-03T17:06:36.341421Z","shell.execute_reply.started":"2023-11-03T17:06:36.310858Z","shell.execute_reply":"2023-11-03T17:06:36.340469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\ndf_s = pd.merge(df_serials, df_injuried,on = 'patient_id')\nX = df_s['serial_id']\ny = df_s[['bowel_healthy',\n          'bowel_injury', \n          'extravasation_healthy', \n          'extravasation_injury',\n          #'kidney_healthy', \n          #'kidney_low', \n          #'kidney_high', \n          'liver_healthy',\n          #'liver_low', \n          #'liver_high', \n          'spleen_healthy',\n          'spleen_low',\n          'spleen_high'\n         ]]\nX_train, X_test, y_train, y_test = train_test_split(\n    X,y , test_size=0.2, stratify=y, random_state=101\n)\n\nX_train, X_val, y_train, y_val = train_test_split(\n    X_train ,y_train , test_size=0.1, stratify=y_train, random_state=101\n)","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:06:36.344877Z","iopub.execute_input":"2023-11-03T17:06:36.345217Z","iopub.status.idle":"2023-11-03T17:06:36.403254Z","shell.execute_reply.started":"2023-11-03T17:06:36.345192Z","shell.execute_reply":"2023-11-03T17:06:36.402461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = df_total[df_total['serial_id'].isin(X_train)]\ntest = df_total[df_total['serial_id'].isin(X_test)]\nval = df_total[df_total['serial_id'].isin(X_val)]","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:06:36.404376Z","iopub.execute_input":"2023-11-03T17:06:36.404698Z","iopub.status.idle":"2023-11-03T17:06:36.430328Z","shell.execute_reply.started":"2023-11-03T17:06:36.404671Z","shell.execute_reply":"2023-11-03T17:06:36.429572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_sequences_with_step_image(data: pd.DataFrame, seq_len, step):\n    sequences = []\n    data_size = len(data)\n\n    for i in tqdm(range(0, data_size - seq_len + 1, step)):\n        feature = data.iloc[i: i + seq_len,3]\n        label_position = (i+seq_len-1)\n        label = data.iloc[label_position,5:-1]\n        \n        sequences.append([feature,label])\n        \n    return sequences\n\nseq_len = num_slices","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:06:36.43134Z","iopub.execute_input":"2023-11-03T17:06:36.431595Z","iopub.status.idle":"2023-11-03T17:06:36.438595Z","shell.execute_reply.started":"2023-11-03T17:06:36.431573Z","shell.execute_reply":"2023-11-03T17:06:36.437752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_s = create_sequences_with_step_image(train,seq_len,seq_len)\ntest_s = create_sequences_with_step_image(test,seq_len,seq_len)\nval_s = create_sequences_with_step_image(val,seq_len,seq_len)","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:06:36.439712Z","iopub.execute_input":"2023-11-03T17:06:36.440036Z","iopub.status.idle":"2023-11-03T17:06:36.771546Z","shell.execute_reply.started":"2023-11-03T17:06:36.440007Z","shell.execute_reply":"2023-11-03T17:06:36.770649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Creation of the test sequences\ntest = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/test_series_meta.csv')\ntest_pat_id = test['patient_id'].unique()\n\n# Define the root directory where the patient data is located\nroot_dir = '/kaggle/input/rsna-2023-abdominal-trauma-detection/test_images'\n\n# Define the list of patient IDs you want to use (replace with your selected IDs)\nselected_patient_ids = test_pat_id \n\n# Initialize empty lists to store data\npatient_ids = []\nimage_paths = []\n\n# Traverse the directories and collect the data\nfor patient_id in selected_patient_ids:\n    patient_dir = os.path.join(root_dir, str(patient_id))\n    if os.path.exists(patient_dir):\n        for serial_id in os.listdir(patient_dir):\n            serial_dir = os.path.join(patient_dir, serial_id)\n            if os.path.isdir(serial_dir):\n                # List all .dcm files in the serial_id folder\n                image_files = [filename for filename in os.listdir(serial_dir) if filename.endswith('.dcm')]\n                \n                # Calculate the step size for evenly spaced selection\n                total_images = len(image_files)\n                if total_images <= num_slices:\n                    step_size = 1  # If fewer than num_slices images, select all of them\n                else:\n                    step_size = total_images // num_slices\n                \n                # Select images in an evenly spaced manner\n                selected_images = [image_files[i] for i in range(0, total_images, step_size)][:num_slices]\n                \n                for filename in selected_images:\n                    image_path = os.path.join(serial_dir, filename)\n                    patient_ids.append(patient_id)\n                    image_paths.append(image_path)\n\n# Create a DataFrame from the collected data\ndf_test = {'patient_id': patient_ids, 'image_path': image_paths}\ndf_test = pd.DataFrame(df_test)\ndf_test_num_slices = pd.DataFrame()\nfor i in df_test.index:\n    df= pd.concat([df_test.iloc[i:i+1,:]]*num_slices,axis=0, ignore_index=True)\n    df_test_num_slices = pd.concat([df_test_num_slices,df],axis=0, ignore_index=True)\ndef create_sequences_with_step_image_inf(data: pd.DataFrame, seq_len, step):\n    sequences = []\n    data_size = len(data)\n\n    for i in tqdm(range(0, data_size - seq_len + 1, step)):\n        feature = data.iloc[i: i + seq_len,1]\n        sequences.append([feature])\n    return sequences\n\nseq_len = num_slices\ntest_seq = create_sequences_with_step_image_inf(df_test_num_slices,seq_len,seq_len)","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:06:36.772828Z","iopub.execute_input":"2023-11-03T17:06:36.773192Z","iopub.status.idle":"2023-11-03T17:06:36.843921Z","shell.execute_reply.started":"2023-11-03T17:06:36.773159Z","shell.execute_reply":"2023-11-03T17:06:36.843068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.save('train_s.npy', train_s)\nnp.save('test_s.npy', train_s)\nnp.save('val_s.npy', train_s)\nnp.save('test_seq.npy', test_seq)","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:21:26.13217Z","iopub.execute_input":"2023-11-03T17:21:26.13254Z","iopub.status.idle":"2023-11-03T17:21:26.138429Z","shell.execute_reply.started":"2023-11-03T17:21:26.132511Z","shell.execute_reply":"2023-11-03T17:21:26.137457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SeqToImageDataset(Dataset):\n\n    def __init__(self,sequences):\n        self.sequences = sequences\n    def __len__(self):\n        return len(self.sequences)\n    def __getitem__(self, idx):   \n        sequence, labels = self.sequences[idx]\n        sequence_features = []\n\n        for image_path in sequence:\n            path = image_path\n            ds = pydicom.dcmread(path)\n            image = ds.pixel_array\n            image = image.astype('int16')\n            image = Image.fromarray(image).convert('RGB')\n            image = image.resize((224, 224))\n            \n            t = transforms.Compose([#transforms.Resize((224, 224)),\n                                    transforms.ToTensor(),\n                                    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])\n            image = t(image)\n            \n        sequence_features.append(image)\n        return {'sequence': sequence_features, 'labels': torch.tensor(labels, dtype=torch.float32)}","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:35:07.635045Z","iopub.execute_input":"2023-11-03T17:35:07.635836Z","iopub.status.idle":"2023-11-03T17:35:07.646818Z","shell.execute_reply.started":"2023-11-03T17:35:07.635782Z","shell.execute_reply":"2023-11-03T17:35:07.645738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SeqToImageDatasetINF(Dataset):\n\n    def __init__(self,sequences):\n        self.sequences = sequences\n    def __len__(self):\n        return len(self.sequences)\n    def __getitem__(self, idx):   \n        sequence = self.sequences[idx]\n        sequence_features = []\n\n        for image_path in sequence[0]:\n            path = image_path\n            ds = pydicom.dcmread(path)\n            image = ds.pixel_array\n            image = image.astype('int16')\n            image = Image.fromarray(image).convert('RGB')\n            image = image.resize((224, 224))\n            \n            t = transforms.Compose([#transforms.Resize((224, 224)),\n                                    transforms.ToTensor(),\n                                    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])\n            image = t(image)\n            \n        sequence_features.append(image)\n        return sequence_features","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:35:07.844595Z","iopub.execute_input":"2023-11-03T17:35:07.845534Z","iopub.status.idle":"2023-11-03T17:35:07.853234Z","shell.execute_reply.started":"2023-11-03T17:35:07.845498Z","shell.execute_reply":"2023-11-03T17:35:07.852338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataModule(L.LightningDataModule):\n    def __init__(self, train_dataframe, val_dataframe, test_dataframe, \n                 pred_dataframe, \n                 batch_size=config.BATCH_SIZE):\n        super().__init__()\n        self.train_dataframe = train_dataframe\n        self.val_dataframe = val_dataframe\n        self.test_dataframe = test_dataframe\n        self.pred_dataframe = pred_dataframe\n\n        self.batch_size = batch_size\n        \n        random.seed(config.SEED)\n        np.random.seed(config.SEED)\n        torch.manual_seed(config.SEED)\n\n    def setup(self, stage=None):\n        self.train_dataset = SeqToImageDataset(self.train_dataframe)\n        self.val_dataset = SeqToImageDataset(self.val_dataframe)\n        self.test_dataset = SeqToImageDataset(self.test_dataframe)\n        self.pred_dataset = SeqToImageDatasetINF(self.pred_dataframe)\n\n    def train_dataloader(self):\n        return DataLoader(\n            dataset=self.train_dataset,\n            batch_size=self.batch_size,\n            shuffle=True,\n            num_workers=1\n        )\n\n    def val_dataloader(self):\n        return DataLoader(\n            dataset=self.val_dataset,\n            batch_size=len(self.val_dataframe)//5,\n            shuffle=False,\n            num_workers=1\n        )\n\n    def test_dataloader(self):\n        return DataLoader(\n            dataset=self.test_dataset,\n            batch_size=len(self.test_dataframe)//5,\n            shuffle=False,\n            num_workers=1\n        )\n\n    def predict_dataloader(self):\n        return DataLoader(\n            dataset=self.pred_dataset,\n            batch_size=self.batch_size,\n            shuffle=False,\n            num_workers=1\n        )","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:35:08.011405Z","iopub.execute_input":"2023-11-03T17:35:08.0118Z","iopub.status.idle":"2023-11-03T17:35:08.0229Z","shell.execute_reply.started":"2023-11-03T17:35:08.01177Z","shell.execute_reply":"2023-11-03T17:35:08.021967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CNNLSTMImage(L.LightningModule):\n    \n    def __init__(self, n_feat=50, n_hidden=150,n_layers=1):\n        super(CNNLSTMImage, self).__init__()\n        self.n_feat = n_feat\n        self.n_hidden = n_hidden\n        self.n_layers = n_layers\n        \n        self.features = torchvision.models.resnet18(weights='DEFAULT')\n        self.features.fc = nn.Linear(in_features=512, out_features=n_feat, bias = True)\n\n\n        self.lstm = nn.LSTM(\n            input_size = n_feat,\n            hidden_size = n_hidden,\n            batch_first = False,\n            num_layers = n_layers\n        )\n        \n        self.layersB = nn.Sequential(\n            nn.GELU(),\n            nn.Linear(in_features=n_hidden, out_features=2, bias=True),\n        )\n        \n        self.layersE = nn.Sequential(\n            nn.GELU(),\n            nn.Linear(in_features=n_hidden, out_features=2, bias=True),\n        )\n        \n        \n        self.layersK = nn.Sequential(\n            nn.GELU(),\n            nn.Linear(in_features=n_hidden, out_features=3, bias=True),\n        )\n        \n        \n        self.layersL = nn.Sequential(\n            nn.GELU(),\n            nn.Linear(in_features=n_hidden, out_features=3, bias=True),\n        )\n  \n        \n        self.layersS = nn.Sequential(\n            nn.GELU(),\n            nn.Linear(in_features=n_hidden, out_features=3, bias=True),\n        )\n        \n    def forward(self, x):\n        \n        sequence_features = []\n\n        for image in x:\n            features = self.features(image)  # Pass each image through the feature extractor\n            sequence_features.append(features)\n\n        sequence_features = torch.stack(sequence_features)\n        self.lstm.flatten_parameters()\n        _, (hidden, _) = self.lstm(sequence_features)\n        out = hidden[-1]\n        \n        bowel_logits = self.layersB(out)\n        extravasation_logits = self.layersE(out)\n        kidney_logits = self.layersK(out)\n        liver_logits = self.layersL(out)\n        spleen_logits = self.layersS(out)\n        \n        return {'bowel':bowel_logits,\n                'extravasation': extravasation_logits,\n                'kidney': kidney_logits,\n                'liver': liver_logits,\n                'spleen': spleen_logits\n               }","metadata":{"execution":{"iopub.status.busy":"2023-11-03T17:35:08.177652Z","iopub.execute_input":"2023-11-03T17:35:08.17845Z","iopub.status.idle":"2023-11-03T17:35:08.191028Z","shell.execute_reply.started":"2023-11-03T17:35:08.178423Z","shell.execute_reply":"2023-11-03T17:35:08.190095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda'\n#Promissor e funcionando\nclass LightM(L.LightningModule):\n    def __init__(self, model, lr=config.LR):\n        super(LightM, self).__init__()\n        self.lr = lr\n        self.model = model\n        \n        self.loss_fn1 = nn.BCEWithLogitsLoss()\n        self.loss_fn2 = nn.BCEWithLogitsLoss()\n        self.loss_fn3 = nn.CrossEntropyLoss()\n        self.loss_fn4 = nn.CrossEntropyLoss()\n        self.loss_fn5 = nn.CrossEntropyLoss()\n\n    def forward(self, x):\n\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        inputs, labels = batch['sequence'], batch['labels']\n\n        # Compute logits for both tasks\n        logits = self(inputs)\n        bowel_logits = logits['bowel']\n        extravasation_logits = logits[\"extravasation\"]\n        kidney_logits = logits['kidney']\n        liver_logits = logits['liver']\n        spleen_logits = logits['spleen']\n\n        # Compute losses and accuracies for both tasks\n        bowel_loss = self.loss_fn1(bowel_logits, labels[:,0:2])\n        extravasation_loss = self.loss_fn2(extravasation_logits, labels[:,2:4])\n        kidney_loss = self.loss_fn3(kidney_logits, labels[:,4:7])\n        liver_loss = self.loss_fn4(liver_logits, labels[:,7:10])\n        spleen_loss = self.loss_fn5(spleen_logits, labels[:,10:])\n        \n        # Compute the total loss (you can weigh the losses if needed)\n        total_loss = bowel_loss + extravasation_loss + kidney_loss + liver_loss + spleen_loss\n\n        self.log('train_bowel_loss', bowel_loss)\n        self.log('train_extravasation_loss', extravasation_loss)\n        self.log('train_kidney_loss', kidney_loss)\n        self.log('train_liver_loss', liver_loss)\n        self.log('train_spleen_loss',spleen_loss)\n        self.log('train_loss',total_loss, prog_bar=True, logger=True)\n        \n\n        #preds\n        bowel_pred = torch.argmax(bowel_logits, dim=1)\n        extravasation_pred = torch.argmax(extravasation_logits, dim=1)\n        kidney_pred = torch.argmax(kidney_logits, dim=1) \n        liver_pred = torch.argmax(liver_logits, dim=1)\n        spleen_pred = torch.argmax(spleen_logits, dim=1)\n        #total_pred = torch.cat([bowel_pred, extravasation_pred, kidney_pred, liver_pred, spleen_pred], dim=1)\n        acc1 =torchmetrics.Accuracy(task='binary').to(device)\n        acc2 =torchmetrics.Accuracy(task='binary').to(device)\n        acc3 =torchmetrics.Accuracy(task='multiclass', num_classes=3).to(device)\n        acc4 =torchmetrics.Accuracy(task='multiclass', num_classes=3).to(device)\n        acc5 =torchmetrics.Accuracy(task='multiclass', num_classes=3).to(device)\n        accuracy1 = acc1(bowel_pred,torch.argmax(labels[:,0:2], dim=1))\n        accuracy2 = acc2(extravasation_pred,torch.argmax(labels[:,2:4], dim=1))\n        accuracy3 = acc3(kidney_pred,torch.argmax(labels[:,4:7], dim=1))\n        accuracy4 = acc4(liver_pred,torch.argmax(labels[:,7:10], dim=1))\n        accuracy5 = acc5(spleen_pred,torch.argmax(labels[:,10:], dim=1))\n        acc = (accuracy1 + accuracy2 + accuracy3 + accuracy4 + accuracy5)/5\n        self.log('train_bowel_acc', accuracy1)\n        self.log('train_extravasation_acc', accuracy2)\n        self.log('train_kidney_acc', accuracy3)\n        self.log('train_liver_acc', accuracy4)\n        self.log('train_spleen_acc',accuracy5)\n        self.log('train_acc',acc, prog_bar=True, logger=True)\n\n        return total_loss\n    \n    def validation_step(self, batch, batch_idx):\n        inputs, labels = batch['sequence'], batch['labels']\n\n        # Compute logits for both tasks\n        logits = self(inputs)\n        bowel_logits = logits['bowel']\n        extravasation_logits = logits[\"extravasation\"]\n        kidney_logits = logits['kidney']\n        liver_logits = logits['liver']\n        spleen_logits = logits['spleen']\n\n        # Compute losses and accuracies for both tasks\n        bowel_loss = self.loss_fn1(bowel_logits, labels[:,0:2])\n        extravasation_loss = self.loss_fn2(extravasation_logits, labels[:,2:4])\n        kidney_loss = self.loss_fn3(kidney_logits, labels[:,4:7])\n        liver_loss = self.loss_fn4(liver_logits, labels[:,7:10])\n        spleen_loss = self.loss_fn5(spleen_logits, labels[:,10:])\n\n\n        # Compute the total loss (you can weigh the losses if needed)\n        total_loss = bowel_loss + extravasation_loss + kidney_loss + liver_loss + spleen_loss\n\n        self.log('val_bowel_loss', bowel_loss)\n        self.log('val_extravasation_loss', extravasation_loss)\n        self.log('val_kidney_loss', kidney_loss)\n        self.log('val_liver_loss', liver_loss)\n        self.log('val_spleen_loss',spleen_loss)\n        self.log('val_loss',total_loss, prog_bar=True, logger=True)\n        \n        #preds\n        bowel_pred = torch.argmax(bowel_logits, dim=1)\n        extravasation_pred = torch.argmax(extravasation_logits, dim=1)\n        kidney_pred = torch.argmax(kidney_logits, dim=1) \n        liver_pred = torch.argmax(liver_logits, dim=1)\n        spleen_pred = torch.argmax(spleen_logits, dim=1)\n        #total_pred = torch.cat([bowel_pred, extravasation_pred, kidney_pred, liver_pred, spleen_pred], dim=1)\n        acc1 =torchmetrics.Accuracy(task='binary').to(device)\n        acc2 =torchmetrics.Accuracy(task='binary').to(device)\n        acc3 =torchmetrics.Accuracy(task='multiclass', num_classes=3).to(device)\n        acc4 =torchmetrics.Accuracy(task='multiclass', num_classes=3).to(device)\n        acc5 =torchmetrics.Accuracy(task='multiclass', num_classes=3).to(device)\n        accuracy1 = acc1(bowel_pred,torch.argmax(labels[:,0:2], dim=1))\n        accuracy2 = acc2(extravasation_pred,torch.argmax(labels[:,2:4], dim=1))\n        accuracy3 = acc3(kidney_pred,torch.argmax(labels[:,4:7], dim=1))\n        accuracy4 = acc4(liver_pred,torch.argmax(labels[:,7:10], dim=1))\n        accuracy5 = acc5(spleen_pred,torch.argmax(labels[:,10:], dim=1))\n        acc = (accuracy1 + accuracy2 + accuracy3 + accuracy4 + accuracy5)/5\n        self.log('val_bowel_acc', accuracy1)\n        self.log('val_extravasation_acc', accuracy2)\n        self.log('val_kidney_acc', accuracy3)\n        self.log('val_liver_acc', accuracy4)\n        self.log('val_spleen_acc',accuracy5)\n        self.log('val_acc',acc, prog_bar=True, logger=True)\n        \n        return total_loss\n    \n    def test_step(self, batch, batch_idx):\n        inputs, labels = batch['sequence'], batch['labels']\n\n        # Compute logits for both tasks\n        logits = self(inputs)\n        bowel_logits = logits['bowel']\n        extravasation_logits = logits[\"extravasation\"]\n        kidney_logits = logits['kidney']\n        liver_logits = logits['liver']\n        spleen_logits = logits['spleen']\n\n        # Compute losses and accuracies for both tasks\n        bowel_loss = self.loss_fn1(bowel_logits, labels[:,0:2])\n        extravasation_loss = self.loss_fn2(extravasation_logits, labels[:,2:4])\n        kidney_loss = self.loss_fn3(kidney_logits, labels[:,4:7])\n        liver_loss = self.loss_fn4(liver_logits, labels[:,7:10])\n        spleen_loss = self.loss_fn5(spleen_logits, labels[:,10:])\n\n\n        # Compute the total loss (you can weigh the losses if needed)\n        total_loss = bowel_loss + extravasation_loss + kidney_loss + liver_loss + spleen_loss\n\n        self.log('test_bowel_loss', bowel_loss)\n        self.log('test_extravasation_loss', extravasation_loss)\n        self.log('test_kidney_loss', kidney_loss)\n        self.log('test_liver_loss', liver_loss)\n        self.log('test_spleen_loss',spleen_loss)\n        self.log('test_loss',total_loss, prog_bar=True, logger=True)\n        \n        #preds\n        bowel_pred = torch.argmax(bowel_logits, dim=1)\n        extravasation_pred = torch.argmax(extravasation_logits, dim=1)\n        kidney_pred = torch.argmax(kidney_logits, dim=1) \n        liver_pred = torch.argmax(liver_logits, dim=1)\n        spleen_pred = torch.argmax(spleen_logits, dim=1)\n        #total_pred = torch.cat([bowel_pred, extravasation_pred, kidney_pred, liver_pred, spleen_pred], dim=1)\n        acc1 =torchmetrics.Accuracy(task='binary').to(device)\n        acc2 =torchmetrics.Accuracy(task='binary').to(device)\n        acc3 =torchmetrics.Accuracy(task='multiclass', num_classes=3).to(device)\n        acc4 =torchmetrics.Accuracy(task='multiclass', num_classes=3).to(device)\n        acc5 =torchmetrics.Accuracy(task='multiclass', num_classes=3).to(device)\n        accuracy1 = acc1(bowel_pred,torch.argmax(labels[:,0:2], dim=1))\n        accuracy2 = acc2(extravasation_pred,torch.argmax(labels[:,2:4], dim=1))\n        accuracy3 = acc3(kidney_pred,torch.argmax(labels[:,4:7], dim=1))\n        accuracy4 = acc4(liver_pred,torch.argmax(labels[:,7:10], dim=1))\n        accuracy5 = acc5(spleen_pred,torch.argmax(labels[:,10:], dim=1))\n        acc = (accuracy1 + accuracy2 + accuracy3 + accuracy4 + accuracy5)/5\n        self.log('test_bowel_acc', accuracy1)\n        self.log('test_extravasation_acc', accuracy2)\n        self.log('test_kidney_acc', accuracy3)\n        self.log('test_liver_acc', accuracy4)\n        self.log('test_spleen_acc',accuracy5)\n        self.log('test_acc',acc, prog_bar=True, logger=True)\n        \n        return total_loss\n    \n    def prediction_step(self, batch, batch_idx):\n        inputs = batch\n\n        logits = self(inputs)\n\n        return logits\n\n#    def configure_optimizers(self):\n#        optimizer = torch.optim.AdamW(self.parameters(), lr=self.lr)\n#        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=4)\n#        return optimizer\n    def configure_optimizers(self):\n        opt = torch.optim.AdamW(self.parameters(), lr=self.lr)\n        sch = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=57, eta_min = 0.0005) # New!\n\n        return {\n            \"optimizer\": opt,\n            \"lr_scheduler\": {\n                \"scheduler\": sch,\n                \"monitor\": \"train_loss\",\n                \"interval\": \"step\", # step means \"batch\" here, default: epoch   # New!\n                \"frequency\": 1, # default\n            },\n        }","metadata":{"execution":{"iopub.status.busy":"2023-11-03T18:00:41.469711Z","iopub.execute_input":"2023-11-03T18:00:41.470103Z","iopub.status.idle":"2023-11-03T18:00:41.518627Z","shell.execute_reply.started":"2023-11-03T18:00:41.470072Z","shell.execute_reply":"2023-11-03T18:00:41.517516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"module = CNNLSTMImage()\nmodel = LightM(model=module, lr=config.LR)\nlogger_v1 = TensorBoardLogger(save_dir=\"/kaggle/working/\", name=\"logs_v1\")\ncallback_v1 = [ModelCheckpoint(save_top_k=1,\n                             verbose=True,\n                             monitor='val_loss',\n                             save_last = True,\n                             mode='min',\n                             filename='best_model',\n                             dirpath = logger_v1.log_dir,\n                              ),\n               EarlyStopping(monitor='val_loss',\n                             patience=3,\n                             mode='min',\n                            )\n              ]\n\ntrainer = L.Trainer(\n    max_epochs=config.MAX_EPOCHS,\n    callbacks = callback_v1,\n    accelerator=\"gpu\",\n    logger=logger_v1,\n    deterministic=True,\n    precision='16-mixed',\n    log_every_n_steps=1\n)\n\ndm = CustomDataModule(train_s,\n                         val_s,\n                         test_s,\n                         test_seq,\n                         batch_size=config.BATCH_SIZE,\n                        )","metadata":{"execution":{"iopub.status.busy":"2023-11-03T18:00:42.015437Z","iopub.execute_input":"2023-11-03T18:00:42.0161Z","iopub.status.idle":"2023-11-03T18:00:42.300164Z","shell.execute_reply.started":"2023-11-03T18:00:42.016058Z","shell.execute_reply":"2023-11-03T18:00:42.299266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(model,datamodule=dm)","metadata":{"execution":{"iopub.status.busy":"2023-11-03T18:00:47.20733Z","iopub.execute_input":"2023-11-03T18:00:47.208327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.test(datamodule=dm,ckpt_path='best', verbose=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}