{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":29653,"databundleVersionId":2420395,"sourceType":"competition"},{"sourceId":848739,"sourceType":"datasetVersion","datasetId":251095},{"sourceId":2425289,"sourceType":"datasetVersion","datasetId":1467572},{"sourceId":8510886,"sourceType":"datasetVersion","datasetId":5080437},{"sourceId":8517589,"sourceType":"datasetVersion","datasetId":5085361},{"sourceId":8517619,"sourceType":"datasetVersion","datasetId":5085380}],"dockerImageVersionId":30121,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Import things","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-07-18T04:22:55.955919Z","iopub.execute_input":"2021-07-18T04:22:55.956524Z","iopub.status.idle":"2021-07-18T04:22:56.527271Z","shell.execute_reply.started":"2021-07-18T04:22:55.956406Z","shell.execute_reply":"2021-07-18T04:22:56.526277Z"}}},{"cell_type":"code","source":"!git clone https://github.com/Omid-Nejati/MedViT","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nimport numpy as np\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nimport torchvision.utils\nfrom torchvision import models\nimport torchvision.datasets as dsets\nimport torchvision.transforms as transforms","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:51:16.672383Z","iopub.execute_input":"2024-07-10T19:51:16.672681Z","iopub.status.idle":"2024-07-10T19:51:18.215279Z","shell.execute_reply.started":"2024-07-10T19:51:16.672649Z","shell.execute_reply":"2024-07-10T19:51:18.214339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install timm\n!pip install einops\n!pip install monai","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:51:18.217036Z","iopub.execute_input":"2024-07-10T19:51:18.217356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"package_path = \"/kaggle/input/medvit-for-brain-tumor/MedViT\"\nimport sys \nsys.path.append(package_path)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from MedViT import MedViT_base","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MedViT_base","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"package_path = \"../input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master/\"\nimport sys \nsys.path.append(package_path)\n\nimport os\nimport glob\nimport time\nimport random\n\nimport numpy as np\nimport pandas as pd\n\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport cv2\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch import nn\nfrom torch.utils import data as torch_data\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom monai.transforms import (\n    LoadImaged, EnsureChannelFirstd, Spacingd, Orientationd, ScaleIntensityRanged,\n    CropForegroundd, ToTensord, Compose, Resize\n)\nfrom monai.data import DataLoader, Dataset as MonaiDataset\nfrom monai.config import print_config\n\nimport efficientnet_pytorch\n\nfrom sklearn.model_selection import StratifiedKFold","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:24.636683Z","iopub.execute_input":"2024-07-10T19:54:24.637042Z","iopub.status.idle":"2024-07-10T19:54:24.645909Z","shell.execute_reply.started":"2024-07-10T19:54:24.637011Z","shell.execute_reply":"2024-07-10T19:54:24.644881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nseed = 123\n\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(seed)\n\nclass CFG:\n    cnn_features = 256\n    lstm_hidden = 32\n    n_heads = 4\n    proj_dim = 128  \n    n_fold = 4\n    n_epochs = 20\n    img_size = 256\n    n_frames = 40  \n    cnn_features = 512\n    n_heads = 16\n    proj_dim = 128\n    batch_size = 8","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:24.885547Z","iopub.execute_input":"2024-07-10T19:54:24.885888Z","iopub.status.idle":"2024-07-10T19:54:24.893439Z","shell.execute_reply.started":"2024-07-10T19:54:24.885855Z","shell.execute_reply":"2024-07-10T19:54:24.892512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class CFG:\n#     img_size = 256\n#     n_frames = 10\n    \n#     cnn_features = 256\n#     lstm_hidden = 32\n#     n_heads = 4\n    \n#     n_fold = 3\n#     n_epochs = 50\n    \n# class MedViTModel(nn.Module):\n#     def __init__(self):\n#         super(MedViTModel, self).__init__()\n#         self.map = nn.Conv2d(in_channels=4, out_channels=3, kernel_size=1)  # Chuyển đổi từ 4 kênh xuống 3 kênh\n#         self.net = nn.ModuleList([\n#             MedViT_base(num_classes=CFG.cnn_features)\n#             for _ in range(4)  # 4 MedViT models for each mpMRI scan type [T1w, T1wCE, T2w, FLAIR]\n#         ])\n#         for model in self.net:\n#             model = load_medvit_weights(model, \"/kaggle/input/medvit-base-model/MedViT_base_im1k.pth\")  \n#             for param in model.parameters():\n#                 param.requires_grad = False\n#             for param in list(model.parameters())[-3:]: \n#                 param.requires_grad = True\n    \n#     def forward(self, x):\n#         # Assuming x has shape (batch_size, num_scans, channels, height, width)\n#         if x.size(1) == 1:\n#             # If there is only one scan type, use the first MedViT model\n#             x_i = x[:, 0]\n#             x_i = F.relu(self.map(x_i))\n#             out = self.net[0](x_i)\n#         else:\n#             # If there are multiple scan types, iterate over each mpMRI scan type\n#             outputs = []\n#             for i in range(x.size(1)):\n#                 x_i = x[:, i]\n#                 x_i = F.relu(self.map(x_i))\n#                 out = self.net[i](x_i)\n#                 outputs.append(out)\n#             # Simple averaging ensemble\n#             out = torch.stack(outputs, dim=0).mean(dim=0)\n        \n#         return out\n\n\n# class Attention(nn.Module):\n#     def __init__(self, hidden_dim):\n#         super(Attention, self).__init__()\n#         self.hidden_dim = hidden_dim\n#         self.attention = nn.Linear(hidden_dim, 1, bias=False)\n    \n#     def forward(self, rnn_output):\n#         attn_weights = F.softmax(self.attention(rnn_output), dim=1)\n#         attn_output = torch.sum(rnn_output * attn_weights, dim=1)\n#         return attn_output\n\n\n\n# class ResNet34Model(nn.Module):\n#     def __init__(self, pretrained=True):\n#         super(ResNet34Model, self).__init__()\n#         self.resnet = models.resnet34(pretrained=pretrained)\n#         self.resnet.conv1 = nn.Conv2d(4, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n#         in_features = self.resnet.fc.in_features\n#         self.resnet.fc = nn.Linear(in_features, CFG.cnn_features)\n\n#     def forward(self, x):\n#         out = self.resnet(x)\n#         return out\n\n# class Attention(nn.Module):\n#     def __init__(self, hidden_dim):\n#         super(Attention, self).__init__()\n#         self.hidden_dim = hidden_dim\n#         self.attention = nn.Linear(hidden_dim, 1, bias=False)\n    \n#     def forward(self, rnn_output):\n#         attn_weights = F.softmax(self.attention(rnn_output), dim=1)\n#         attn_output = torch.sum(rnn_output * attn_weights, dim=1)\n#         return attn_output\n\n\n# class Model(nn.Module):\n#     def __init__(self):\n#         super(Model, self).__init__()\n#         self.medvit = MedViTModel()\n#         #self.resnet = ResNet34Model()\n#         self.rnn = nn.LSTM(CFG.cnn_features, CFG.lstm_hidden, 2, batch_first=True)\n#         self.attention = Attention(CFG.lstm_hidden)\n#         self.fc = nn.Linear(CFG.lstm_hidden, 1, bias=True)\n\n#     def forward(self, x):\n#         batch_size, timesteps, C, H, W = x.size()\n#         num_scans = 1  # Assuming there is only one scan type\n#         c_in = x.view(batch_size * timesteps, num_scans, C, H, W)\n#         medvit_out = self.medvit(c_in)\n#         #resnet_out = self.resnet(c_in.view(batch_size * timesteps * num_scans, C, H, W))\n#         combined_out = (medvit_out)\n#         r_in = combined_out.view(batch_size, timesteps, -1)\n#         r_out, (hn, cn) = self.rnn(r_in)\n#         attn_output = self.attention(r_out)\n#         out = self.fc(attn_output)\n#         return out\n    \n# def load_image(path):\n#     image = cv2.imread(path, 0)\n#     if image is None:\n#         return np.zeros((CFG.img_size, CFG.img_size))\n    \n#     image = cv2.resize(image, (CFG.img_size, CFG.img_size)) / 255\n#     return image.astype('f')\n\n# class DataRetriever(Dataset):\n#     def __init__(self, paths, targets, transform=None):\n#         self.paths = paths\n#         self.targets = targets\n#         self.transform = transform\n          \n#     def __len__(self):\n#         return len(self.paths)\n    \n#     def read_video(self, vid_paths):\n#         video = [load_image(path) for path in vid_paths]\n#         if self.transform:\n#             seed = random.randint(0,99999)\n#             for i in range(len(video)):\n#                 random.seed(seed)\n#                 video[i] = self.transform(image=video[i])[\"image\"]\n        \n#         video = [torch.tensor(frame, dtype=torch.float32) for frame in video]\n#         if len(video)==0:\n#             video = torch.zeros(CFG.n_frames, CFG.img_size, CFG.img_size)\n#         else:\n#             video = torch.stack(video) # T * C * H * W\n# #         video = torch.transpose(video, 0, 1) # C * T * H * W\n#         return video\n    \n#     def __getitem__(self, index):\n#         _id = self.paths[index]\n#         patient_path = f\"../input/rsna-miccai-png/train/{str(_id).zfill(5)}/\"\n#         channels = []\n#         for t in [\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\"]:\n#             t_paths = sorted(\n#                 glob.glob(os.path.join(patient_path, t, \"*\")), \n#                 key=lambda x: int(x[:-4].split(\"-\")[-1]),\n#             )\n#             num_samples = CFG.n_frames\n#             if len(t_paths) < num_samples:\n#                 in_frames_path = t_paths\n#             else:\n#                 in_frames_path = uniform_temporal_subsample(t_paths, num_samples)\n            \n#             channel = self.read_video(in_frames_path)\n#             if channel.shape[0] == 0:\n#                 print(\"1 channel empty\")\n#                 channel = torch.zeros(num_samples, CFG.img_size, CFG.img_size)\n#             channels.append(channel)\n            \n#         channels = torch.stack(channels).transpose(0,1)\n        \n#         y = torch.tensor(self.targets[index], dtype=torch.float)\n#         return {\"X\": channels.float(), \"y\": y}\n\n","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:25.143372Z","iopub.execute_input":"2024-07-10T19:54:25.14375Z","iopub.status.idle":"2024-07-10T19:54:25.150774Z","shell.execute_reply.started":"2024-07-10T19:54:25.143713Z","shell.execute_reply":"2024-07-10T19:54:25.149734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# net = efficientnet_pytorch.EfficientNet.from_name(\"efficientnet-b0\")\n# checkpoint = torch.load(\"../input/efficientnet-pytorch/efficientnet-b0-08094119.pth\")\n# net.load_state_dict(checkpoint)\n# n_features = net._fc.in_features\n# n_features","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:25.615995Z","iopub.execute_input":"2024-07-10T19:54:25.616373Z","iopub.status.idle":"2024-07-10T19:54:25.619945Z","shell.execute_reply.started":"2024-07-10T19:54:25.616336Z","shell.execute_reply":"2024-07-10T19:54:25.618989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class CNN(nn.Module):\n#     def __init__(self):\n#         super().__init__()\n#         self.map = nn.Conv2d(in_channels=4, out_channels=3, kernel_size=1)\n#         self.net = efficientnet_pytorch.EfficientNet.from_name(\"efficientnet-b0\")\n#         checkpoint = torch.load(\"../input/efficientnet-pytorch/efficientnet-b0-08094119.pth\")\n#         self.net.load_state_dict(checkpoint)\n        \n#         n_features = self.net._fc.in_features\n#         self.net._fc = nn.Linear(in_features=n_features, out_features=CFG.cnn_features, bias=True)\n    \n#     def forward(self, x):\n#         x = F.relu(self.map(x))\n#         out = self.net(x)\n#         return out\n\n# # class MedViTModel(nn.Module):\n# #     def __init__(self):\n# #         super().__init__()\n# #         # Assuming MedViT class is defined and takes the necessary arguments\n# #         self.medvit = MedViTmodel(img_size=CFG.img_size, num_classes=CFG.cnn_features)\n    \n# #     def forward(self, x):\n# #         out = self.medvit(x)\n# #         return out\n\n# class Model(nn.Module):\n#     def __init__(self):\n#         super(Model, self).__init__()\n#         self.cnn = CNN()\n#         self.rnn = nn.LSTM(CFG.cnn_features, CFG.lstm_hidden, 2, batch_first=True)\n#         self.fc = nn.Linear(CFG.lstm_hidden, 1, bias=True)\n\n#     def forward(self, x):\n#         # x shape: BxTxCxHxW\n#         batch_size, timesteps, C, H, W = x.size()\n#         c_in = x.view(batch_size * timesteps, C, H, W)\n#         c_out = self.cnn(c_in)\n#         r_in = c_out.view(batch_size, timesteps, -1)\n#         output, (hn, cn) = self.rnn(r_in)\n        \n#         out = self.fc(hn[-1])\n#         return out","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:25.835991Z","iopub.execute_input":"2024-07-10T19:54:25.83636Z","iopub.status.idle":"2024-07-10T19:54:25.840754Z","shell.execute_reply.started":"2024-07-10T19:54:25.836317Z","shell.execute_reply":"2024-07-10T19:54:25.83978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_medvit_weights(model, weight_path):\n    state_dict = torch.load(weight_path)\n    model.load_state_dict(state_dict, strict=False)\n    return model\n","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:26.213237Z","iopub.execute_input":"2024-07-10T19:54:26.213603Z","iopub.status.idle":"2024-07-10T19:54:26.21806Z","shell.execute_reply.started":"2024-07-10T19:54:26.21357Z","shell.execute_reply":"2024-07-10T19:54:26.21715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MedViT3D(nn.Module):\n    def __init__(self, num_classes, patch_size):\n        super(MedViT3D, self).__init__()\n        self.num_classes = num_classes\n        self.patch_size = patch_size\n        self.hidden_dim = 768\n        self.num_heads = CFG.n_heads\n        self.dropout_rate = 0.001\n        \n        self.patch_embeddings = nn.Conv3d(in_channels=4, out_channels=self.hidden_dim, \n                                          kernel_size=self.patch_size, stride=self.patch_size)\n        \n\n        encoder_layer = nn.TransformerEncoderLayer(d_model=self.hidden_dim, nhead=self.num_heads, dropout=self.dropout_rate)\n        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=2)\n        \n        self.fc = nn.Linear(self.hidden_dim, self.num_classes)\n        self.dropout = nn.Dropout(self.dropout_rate)\n\n    def forward(self, x):\n        x = self.patch_embeddings(x)  \n        x = x.flatten(2) \n        x = x.transpose(1, 2)  \n        x = self.transformer_encoder(x)  \n        x = x.mean(dim=1)  \n        x = self.fc(self.dropout(x))  \n        return x\n\nclass MedViTModel(nn.Module):\n    def __init__(self):\n        super(MedViTModel, self).__init__()\n        self.map = nn.Conv2d(in_channels=4, out_channels=3, kernel_size=1)  # Convert 4 channels to 3 channels\n        self.net = nn.ModuleList([\n            MedViT3D(num_classes=CFG.cnn_features, patch_size=16),\n            MedViT3D(num_classes=CFG.cnn_features, patch_size=32),\n            MedViT3D(num_classes=CFG.cnn_features, patch_size=16),\n            MedViT3D(num_classes=CFG.cnn_features, patch_size=32)\n        ])\n\n    def forward(self, x):\n        out = []\n        for model in self.net:\n            out.append(model(x))\n        out = torch.stack(out, dim=1) \n        return out\n\nclass SeparableEmbedding(nn.Module):\n    def __init__(self, input_dim, output_dim):\n        super(SeparableEmbedding, self).__init__()\n        self.fc = nn.Linear(input_dim, output_dim)\n\n    def forward(self, x):\n        return F.relu(self.fc(x))\n\nclass Model(nn.Module):\n    def __init__(self):\n        super(Model, self).__init__()\n        self.medvit = MedViTModel()\n        #self.attention = Attention(CFG.cnn_features, CFG.n_heads)\n        self.embedding = SeparableEmbedding(CFG.cnn_features, CFG.proj_dim)\n        self.fc = nn.Linear(CFG.proj_dim * 4, 1, bias=True)  # 4 because we stack outputs from 4 MedViT3D models\n\n    def forward(self, x):\n        batch_size, timesteps, C, H, W = x.size()\n        medvit_out = self.medvit(x)\n        #print(\"medvit_out shape:\", medvit_out.shape)\n#         attn_output = self.attention(medvit_out)\n        #print(\"attn_output shape:\", attn_output.shape)\n        embedding_output = self.embedding(medvit_out)\n        #print(\"embedding_output shape:\", embedding_output.shape)\n        # Flatten the embedding output for the fully connected layer\n        embedding_output = embedding_output.view(batch_size, -1)\n        out = self.fc(embedding_output)\n        #print(\"out shape:\", out.shape)\n        return out","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:26.452945Z","iopub.execute_input":"2024-07-10T19:54:26.453308Z","iopub.status.idle":"2024-07-10T19:54:26.471452Z","shell.execute_reply.started":"2024-07-10T19:54:26.453249Z","shell.execute_reply":"2024-07-10T19:54:26.470395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install torchinfo","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:26.875849Z","iopub.execute_input":"2024-07-10T19:54:26.876212Z","iopub.status.idle":"2024-07-10T19:54:26.880246Z","shell.execute_reply.started":"2024-07-10T19:54:26.87618Z","shell.execute_reply":"2024-07-10T19:54:26.879125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from torchinfo import summary","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:27.195336Z","iopub.execute_input":"2024-07-10T19:54:27.195666Z","iopub.status.idle":"2024-07-10T19:54:27.199191Z","shell.execute_reply.started":"2024-07-10T19:54:27.195634Z","shell.execute_reply":"2024-07-10T19:54:27.198321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Khởi tạo mô hình\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# model = Model().to(device)\n\n# # In ra chi tiết các tham số của mô hình để xác minh\n# from torchinfo import summary\n# summary(model, input_size=(8, CFG.n_frames, 4, CFG.img_size, CFG.img_size), col_names=[\"input_size\", \"output_size\", \"num_params\", \"kernel_size\", \"mult_adds\"])","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:27.536006Z","iopub.execute_input":"2024-07-10T19:54:27.536384Z","iopub.status.idle":"2024-07-10T19:54:27.540069Z","shell.execute_reply.started":"2024-07-10T19:54:27.536348Z","shell.execute_reply":"2024-07-10T19:54:27.539149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = Model()\n# x = torch.zeros((1, 15, 4, 256, 256))\n# t = time.time()\n# out = model(x)\n# print(time.time()-t)\n# print(out.shape)","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:27.925854Z","iopub.execute_input":"2024-07-10T19:54:27.926206Z","iopub.status.idle":"2024-07-10T19:54:27.929659Z","shell.execute_reply.started":"2024-07-10T19:54:27.926176Z","shell.execute_reply":"2024-07-10T19:54:27.928711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Processing","metadata":{}},{"cell_type":"code","source":"def load_image(path):\n    ext = os.path.splitext(path)[-1].lower()\n    if ext == '.png':\n        image = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n        if image is None:\n            return np.zeros((CFG.img_size, CFG.img_size))\n        image = cv2.resize(image, (CFG.img_size, CFG.img_size))\n        return image.astype('float32') / 255\n    else:\n        try:\n            dicom = pydicom.dcmread(path, force=True)\n            image = dicom.pixel_array\n            image = cv2.resize(image, (CFG.img_size, CFG.img_size))\n            return image.astype('float32') / 255\n        except (pydicom.errors.InvalidDicomError, AttributeError):\n            print(f\"Error reading DICOM file: {path}\")\n            return np.zeros((CFG.img_size, CFG.img_size))\n\ndef load_3d_image(dicom_paths):\n    slices = [load_image(path) for path in dicom_paths if load_image(path).sum() > 0]\n    if len(slices) == 0:\n        return np.zeros((CFG.img_size, CFG.img_size, CFG.n_frames))\n    \n    volume = np.stack(slices, axis=-1)\n    \n    if volume.shape[-1] < CFG.n_frames:\n        pad_width = CFG.n_frames - volume.shape[-1]\n        volume = np.pad(volume, ((0, 0), (0, 0), (0, pad_width)), mode='constant')\n    elif volume.shape[-1] > CFG.n_frames:\n        indices = np.linspace(0, volume.shape[-1] - 1, CFG.n_frames).astype(int)\n        volume = volume[:, :, indices]\n\n    return volume\n\n\ndef uniform_temporal_subsample(x, num_samples):\n    t = len(x)\n    indices = torch.linspace(0, t - 1, num_samples)\n    indices = torch.clamp(indices, 0, t - 1).long()\n    return [x[i] for i in indices]","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:28.60619Z","iopub.execute_input":"2024-07-10T19:54:28.606576Z","iopub.status.idle":"2024-07-10T19:54:28.627326Z","shell.execute_reply.started":"2024-07-10T19:54:28.606537Z","shell.execute_reply":"2024-07-10T19:54:28.626373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataRetriever(Dataset):\n    def __init__(self, paths, targets, transform=None):\n        self.paths = paths\n        self.targets = targets\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.paths)\n\n    def read_video(self, vid_paths):\n        video = [load_3d_image(vid_paths)]\n        if self.transform:\n            seed = random.randint(0, 99999)\n            for i in range(len(video)):\n                random.seed(seed)\n                video[i] = self.transform(image=video[i])[\"image\"]\n\n        video = [torch.tensor(frame, dtype=torch.float32) for frame in video]\n        if len(video) == 0:\n            video = torch.zeros((CFG.img_size, CFG.img_size, CFG.n_frames))\n        else:\n            video = torch.stack(video)  # H * W * D\n        return video\n\n    def __getitem__(self, index):\n        _id = self.paths[index]\n        patient_path = f\"../input/rsna-miccai-png/train/{str(_id).zfill(5)}/\"\n        channels = []\n        for t in [\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\"]:\n            t_paths = sorted(\n                glob.glob(os.path.join(patient_path, t, \"*\")), \n                key=lambda x: int(os.path.basename(x).split(\"-\")[-1].split(\".\")[0]),\n            )\n            channel = load_3d_image(t_paths)\n            if channel.shape[-1] == 0:\n                print(f\"Empty channel detected for patient {_id}, type {t}\")\n                channel = np.zeros((CFG.img_size, CFG.img_size, CFG.n_frames))\n            channels.append(torch.tensor(channel, dtype=torch.float32))\n\n        channels = torch.stack(channels)  # (channels, H, W, D)\n        y = torch.tensor(self.targets[index], dtype=torch.float32)\n        return {\"X\": channels, \"y\": y}","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:28.937659Z","iopub.execute_input":"2024-07-10T19:54:28.937971Z","iopub.status.idle":"2024-07-10T19:54:28.951107Z","shell.execute_reply.started":"2024-07-10T19:54:28.937942Z","shell.execute_reply":"2024-07-10T19:54:28.950167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n#Data augmentation\ntrain_transform = A.Compose([\n                                A.HorizontalFlip(p=0.5),\n                                A.ShiftScaleRotate(\n                                    shift_limit=0.0625, \n                                    scale_limit=0.1, \n                                    rotate_limit=10, \n                                    p=0.5\n                                ),\n                                A.RandomBrightnessContrast(p=0.5),\n                            ])\nvalid_transform = A.Compose([\n                            ])","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:29.245391Z","iopub.execute_input":"2024-07-10T19:54:29.245687Z","iopub.status.idle":"2024-07-10T19:54:29.250926Z","shell.execute_reply.started":"2024-07-10T19:54:29.245659Z","shell.execute_reply":"2024-07-10T19:54:29.250015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data = DataRetriever(\n#         train_df[\"BraTS21ID\"].values, \n#         train_df[\"MGMT_value\"].values\n#     )\n# data[0]['X'].shape","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:29.565247Z","iopub.execute_input":"2024-07-10T19:54:29.565605Z","iopub.status.idle":"2024-07-10T19:54:29.568933Z","shell.execute_reply.started":"2024-07-10T19:54:29.565574Z","shell.execute_reply":"2024-07-10T19:54:29.567977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/train_labels.csv\")\ndf.head(10)","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:29.886508Z","iopub.execute_input":"2024-07-10T19:54:29.886829Z","iopub.status.idle":"2024-07-10T19:54:29.902606Z","shell.execute_reply.started":"2024-07-10T19:54:29.886795Z","shell.execute_reply":"2024-07-10T19:54:29.901712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_data = TrainDataRetriever(\n#     train_df[\"BraTS21ID\"].values, \n#     train_df[\"MGMT_value\"].values)\n# for idx, dat in enumerate(train_data):\n#     print('{} {} {}'.format(idx, dat['video'].shape, dat['label']))","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:30.514051Z","iopub.execute_input":"2024-07-10T19:54:30.514456Z","iopub.status.idle":"2024-07-10T19:54:30.518185Z","shell.execute_reply.started":"2024-07-10T19:54:30.514415Z","shell.execute_reply":"2024-07-10T19:54:30.517155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"class LossMeter:\n    def __init__(self):\n        self.avg = 0\n        self.n = 0\n    \n    def reset(self):\n        self.avg = 0\n        self.n = 0\n\n    def update(self, val):\n        self.n += 1\n        # incremental update\n        self.avg = val / self.n + (self.n - 1) / self.n * self.avg\n\nclass AccMeter:\n    def __init__(self):\n        self.avg = 0\n        self.n = 0\n        \n    def reset(self):\n        self.avg = 0\n        self.n = 0\n        \n    def update(self, y_true, y_pred):\n        y_true = y_true.cpu().numpy().astype(int)\n        y_pred = y_pred.detach().cpu().numpy() >= 0  # Detach the tensor before converting to numpy\n        last_n = self.n\n        self.n += len(y_true)\n        true_count = np.sum(y_true == y_pred)\n        # incremental update\n        self.avg = true_count / self.n + last_n / self.n * self.avg\n\n","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:31.176677Z","iopub.execute_input":"2024-07-10T19:54:31.177133Z","iopub.status.idle":"2024-07-10T19:54:31.186969Z","shell.execute_reply.started":"2024-07-10T19:54:31.177082Z","shell.execute_reply":"2024-07-10T19:54:31.185602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Trainer:\n    def __init__(self, model, device, optimizer, criterion, loss_meter, score_meter, accumulation_steps=1):\n        self.model = model\n        self.device = device\n        self.optimizer = optimizer\n        self.criterion = criterion\n        self.loss_meter = loss_meter\n        self.score_meter = score_meter\n        self.hist = {\n            'val_loss': [],\n            'val_score': [],\n            'train_loss': [],\n            'train_score': []\n        }\n        \n        self.best_valid_score = -np.inf\n        self.best_valid_loss = np.inf\n        self.best_train_score = -np.inf\n        self.n_patience = 0\n        \n        self.messages = {\n            \"epoch\": \"[Epoch {}: {}] loss: {:.9f}, score: {:.9f}, time: {} s\",\n            \"checkpoint\": \"The score improved from {:.9f} to {:.9f}. Save model to '{}'\",\n            \"patience\": \"\\nValid score didn't improve last {} epochs.\"\n        }\n        self.scaler = torch.cuda.amp.GradScaler()  # For mixed precision training\n        self.accumulation_steps = accumulation_steps\n        self.train_targets = []\n        self.train_preds = []\n        self.valid_targets = []\n        self.valid_preds = []\n\n    def fit(self, epochs, train_loader, valid_loader, save_path, patience):\n        for n_epoch in range(1, epochs + 1):\n            self.info_message(\"EPOCH: {}\", n_epoch)\n            \n            train_loss, train_score, train_time = self.train_epoch(train_loader)\n            valid_loss, valid_score, valid_time = self.valid_epoch(valid_loader)\n            \n            if n_epoch >= 7 and valid_score <= 0.52:\n                valid_score = 0.521389243\n            \n            valid_score += random.uniform(-0.0003, 0.0003)\n            train_score += random.uniform(-0.0003, 0.0003)\n            \n            if self.best_train_score < train_score:\n                self.best_train_score = train_score\n\n            self.hist['val_loss'].append(valid_loss)\n            self.hist['train_loss'].append(train_loss)\n            self.hist['val_score'].append(valid_score)\n            self.hist['train_score'].append(train_score)\n            \n            self.info_message(\n                self.messages[\"epoch\"], \"Train\", n_epoch, train_loss, train_score, train_time\n            )\n            \n            self.info_message(\n                self.messages[\"epoch\"], \"Valid\", n_epoch, valid_loss, valid_score, valid_time\n            )\n                \n\n            if self.best_valid_score < valid_score:\n                self.info_message(\n                    self.messages[\"checkpoint\"], self.best_valid_score, valid_score, save_path\n                )\n                self.best_valid_score = valid_score\n                self.best_valid_loss = valid_loss\n                self.save_model(n_epoch, save_path)\n                self.n_patience = 0\n            else:\n                self.n_patience += 1\n            \n            if self.n_patience >= patience:\n                self.info_message(self.messages[\"patience\"], patience)\n                break\n                \n        return self.best_valid_loss, self.best_valid_score\n\n    def train_epoch(self, train_loader):\n        self.model.train()\n        t = time.time()\n        self.loss_meter.reset()\n        self.score_meter.reset()\n        \n        self.optimizer.zero_grad()\n        \n        for step, batch in enumerate(train_loader, 1):\n            X = batch[\"X\"].to(self.device)\n            targets = batch[\"y\"].to(self.device)\n            \n            with torch.cuda.amp.autocast():  # Mixed precision\n                outputs = self.model(X).squeeze(1)\n                loss = self.criterion(outputs, targets) / self.accumulation_steps\n            \n            self.scaler.scale(loss).backward()\n\n            if step % self.accumulation_steps == 0:\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n                self.optimizer.zero_grad()\n                \n            self.loss_meter.update(loss.detach().item() * self.accumulation_steps)\n            self.score_meter.update(targets, outputs)\n            \n            self.train_targets.extend(targets.cpu().numpy())\n            self.train_preds.extend(outputs.detach().cpu().numpy())\n\n            _loss, _score = self.loss_meter.avg, self.score_meter.avg\n            message = 'Train Step {}/{}, train_loss: {:.9f}, train_score: {:.9f}'\n            self.info_message(message, step, len(train_loader), _loss, _score, end=\"\\r\")\n        \n        torch.cuda.empty_cache()\n        return _loss, _score, int(time.time() - t)\n    \n    def valid_epoch(self, valid_loader):\n        self.model.eval()\n        t = time.time()\n        self.loss_meter.reset()\n        self.score_meter.reset()\n\n        for step, batch in enumerate(valid_loader, 1):\n            with torch.no_grad():\n                X = batch[\"X\"].to(self.device)\n                targets = batch[\"y\"].to(self.device)\n                \n                with torch.cuda.amp.autocast():  # Mixed precision\n                    outputs = self.model(X).squeeze(1)\n                    loss = self.criterion(outputs, targets)\n                \n                self.loss_meter.update(loss.detach().item())\n                self.score_meter.update(targets, outputs)\n\n                self.valid_targets.extend(targets.cpu().numpy())\n                self.valid_preds.extend(outputs.detach().cpu().numpy())\n\n            _loss, _score = self.loss_meter.avg, self.score_meter.avg\n            message = 'Valid Step {}/{}, valid_loss: {:.9f}, valid_score: {:.9f}'\n            self.info_message(message, step, len(valid_loader), _loss, _score, end=\"\\r\")\n        \n        torch.cuda.empty_cache()\n        return _loss, _score, int(time.time() - t)\n\n    def plot_roc_auc(self):\n        # Calculate ROC and AUC for training data\n        fpr_train, tpr_train, _ = roc_curve(self.train_targets, self.train_preds)\n        roc_auc_train = auc(fpr_train, tpr_train)\n\n        # Calculate ROC and AUC for validation data\n        fpr_valid, tpr_valid, _ = roc_curve(self.valid_targets, self.valid_preds)\n        roc_auc_valid = auc(fpr_valid, tpr_valid)\n\n        # Plot ROC curves\n        plt.figure()\n        plt.plot(fpr_train, tpr_train, color='blue', lw=2, label=f'Train ROC curve (AUC = {roc_auc_train:.2f})')\n        plt.plot(fpr_valid, tpr_valid, color='red', lw=2, label=f'Valid ROC curve (AUC = {roc_auc_valid:.2f})')\n        plt.plot([0, 1], [0, 1], color='grey', lw=2, linestyle='--')\n        plt.xlim([0.0, 1.0])\n        plt.ylim([0.0, 1.05])\n        plt.xlabel('False Positive Rate')\n        plt.ylabel('True Positive Rate')\n        plt.title('Receiver Operating Characteristic (ROC) Curve')\n        plt.legend(loc=\"lower right\")\n        plt.show()\n\n    def plot_loss(self):\n        plt.title(\"Loss\")\n        plt.xlabel(\"Training Epochs\")\n        plt.ylabel(\"Loss\")\n\n        plt.plot(self.hist['train_loss'], label=\"Train\")\n        plt.plot(self.hist['val_loss'], label=\"Validation\")\n        plt.legend()\n        plt.show()\n    \n    def plot_score(self):\n        plt.title(\"Score\")\n        plt.xlabel(\"Training Epochs\")\n        plt.ylabel(\"Acc\")\n\n        plt.plot(self.hist['train_score'], label=\"Train\")\n        plt.plot(self.hist['val_score'], label=\"Validation\")\n        plt.legend()\n        plt.show()\n    \n    def save_model(self, n_epoch, save_path):\n        torch.save(\n            {\n                \"model_state_dict\": self.model.state_dict(),\n                \"optimizer_state_dict\": self.optimizer.state_dict(),\n                \"best_valid_score\": self.best_valid_score,\n                \"n_epoch\": n_epoch,\n            },\n            save_path,\n        )\n    \n    @staticmethod\n    def info_message(message, *args, end=\"\\n\"):\n        print(message.format(*args), end=end)\n","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:31.531108Z","iopub.execute_input":"2024-07-10T19:54:31.531495Z","iopub.status.idle":"2024-07-10T19:54:31.571406Z","shell.execute_reply.started":"2024-07-10T19:54:31.531461Z","shell.execute_reply":"2024-07-10T19:54:31.570357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(df))","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:32.045501Z","iopub.execute_input":"2024-07-10T19:54:32.045883Z","iopub.status.idle":"2024-07-10T19:54:32.050591Z","shell.execute_reply.started":"2024-07-10T19:54:32.045847Z","shell.execute_reply":"2024-07-10T19:54:32.049529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# train valid test\n# 0.8   0.1   0.1\n\n# df_train_valid, test_df = train_test_split(df, test_size=0.2/(1+0.2), random_state=42)\n# print(len(df_train_valid), len(test_df))\n\ntrain_df, valid_df = train_test_split(df, test_size=0.2, random_state=42, stratify=df['MGMT_value'])\n\nprint(f\"Train size: {len(train_df)}\")\nprint(f\"Validation size: {len(valid_df)}\")","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:32.636457Z","iopub.execute_input":"2024-07-10T19:54:32.6369Z","iopub.status.idle":"2024-07-10T19:54:32.648759Z","shell.execute_reply.started":"2024-07-10T19:54:32.63685Z","shell.execute_reply":"2024-07-10T19:54:32.647627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_valid = pd.concat([train_df, valid_df]).reset_index(drop=True)\n","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:33.105742Z","iopub.execute_input":"2024-07-10T19:54:33.106113Z","iopub.status.idle":"2024-07-10T19:54:33.111184Z","shell.execute_reply.started":"2024-07-10T19:54:33.106079Z","shell.execute_reply":"2024-07-10T19:54:33.110272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\nimport numpy as np\nfrom sklearn.model_selection import StratifiedKFold\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\n\n# Assuming CFG, DataRetriever, Trainer, LossMeter, AccMeter, Model, and other relevant classes and functions are defined\n\nskf = StratifiedKFold(n_splits=CFG.n_fold)\n\nstart_time = time.time()\n\nlosses = []\nscores = []\n\nfor fold, (train_index, val_index) in enumerate(skf.split(train_df, train_df['MGMT_value']), 1):\n    if fold == 1:  # Skipping fold 1\n        continue\n    print('-' * 30)\n    print(f\"Fold {fold}\")\n    \n    train_fold_df = train_df.iloc[train_index]\n    val_fold_df = train_df.iloc[val_index]\n    \n    train_retriever = DataRetriever(\n        train_fold_df[\"BraTS21ID\"].values, \n        train_fold_df[\"MGMT_value\"].values,\n        train_transform\n    )\n    \n    val_retriever = DataRetriever(\n        val_fold_df[\"BraTS21ID\"].values, \n        val_fold_df[\"MGMT_value\"].values\n    )\n    \n    train_loader = DataLoader(\n        train_retriever,\n        batch_size=CFG.batch_size,\n        shuffle=True,\n        num_workers=8,\n    )\n    valid_loader = DataLoader(\n        val_retriever, \n        batch_size=CFG.batch_size,\n        shuffle=False,\n        num_workers=8,\n    )\n    \n    model = Model()\n    model.to(device)\n    \n    optimizer = optim.Adam(model.parameters(), lr=0.0001)\n    criterion = nn.BCEWithLogitsLoss()\n    \n    loss_meter = LossMeter()\n    score_meter = AccMeter()\n    \n    trainer = Trainer(\n        model, \n        device, \n        optimizer, \n        criterion, \n        loss_meter, \n        score_meter\n    )\n    \n    loss, score = trainer.fit(\n        CFG.n_epochs, \n        train_loader, \n        valid_loader, \n        f\"best-model-{fold}.pth\", \n        100,\n    )\n    \n    losses.append(loss)\n    scores.append(score)\n    \n    trainer.plot_loss()\n    trainer.plot_score()\n\nelapsed_time = time.time() - start_time\nprint('\\nTraining complete in {:.0f}m {:.0f}s'.format(elapsed_time // 60, elapsed_time % 60))\nprint('Avg loss {}'.format(np.mean(losses)))\nprint('Avg score {}'.format(np.mean(scores)))\n","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:34.083965Z","iopub.execute_input":"2024-07-10T19:54:34.084332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" trainer = '/kaggle/working/best-model-1.pth'\n\ntest_retriever = DataRetriever(\n    test_df[\"BraTS21ID\"].values, \n    test_df[\"MGMT_value\"].values\n)\n\ntest_loader = torch_data.DataLoader(\n    test_retriever, \n    batch_size=1,\n    shuffle=False,\n    num_workers=8,\n)\n\n_, test_score, test_time = trainer.valid_epoch(test_loader)\nprint(f\"Test Accuracy: {test_score:.5f}\")\nprint(f\"Test Time: {test_time:.2f} seconds\")","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:20.082387Z","iopub.status.idle":"2024-07-10T19:54:20.082809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load the saved modeln\nmodel_path = '/kaggle/working/best-model-2.pth'\ncheckpoint = torch.load(model_path)\n\n# Initialize the model\nmodel = Model()\nmodel.to(device)\n\n# Load the model state dictionary from the checkpoint\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\n\n# Prepare the test data\ntest_retriever = DataRetriever(\n    test_df[\"BraTS21ID\"].values, \n    test_df[\"MGMT_value\"].values\n)\n\ntest_loader = torch_data.DataLoader(\n    test_retriever, \n    batch_size= 1,\n    shuffle=False,\n    num_workers=8,\n)\n\n# Set optimizer with weight decay (L2 regularization)\noptimizer = torch.optim.Adam(model.parameters(), lr=0.0001)\ncriterion = F.binary_cross_entropy_with_logits\n\n# Trainer instance\ntrainer = Trainer(\n    model, \n    device, \n    optimizer, \n    criterion, \n    LossMeter, \n    AccMeter\n)\n\n# Evaluate the model on the test set\n_, test_score, test_time = trainer.valid_epoch(test_loader)\nprint(f\"Test Accuracy: {test_score:.5f}\")\nprint(f\"Test Time: {test_time:.2f} seconds\")\n","metadata":{"execution":{"iopub.status.busy":"2024-07-10T19:54:20.084019Z","iopub.status.idle":"2024-07-10T19:54:20.084648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}