{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# IMPORTS","metadata":{}},{"cell_type":"code","source":"package_path =  \"../input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master/\"\nimport sys\nsys.path.append(package_path)\n\nfrom glob import glob\nimport time\nimport random \nimport os\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 \nimport torch.nn as nn\nfrom torch.utils import data as torch_data\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport efficientnet_pytorch\nfrom torchvision import transforms\nfrom PIL import Image\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.model_selection import train_test_split\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-16T09:48:58.921329Z","iopub.execute_input":"2023-09-16T09:48:58.921862Z","iopub.status.idle":"2023-09-16T09:49:05.140746Z","shell.execute_reply.started":"2023-09-16T09:48:58.921826Z","shell.execute_reply":"2023-09-16T09:49:05.139508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configure","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available else \"cpu\")\nseed = 123\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\nseed_everything(seed)\n\nclass CFG:\n    img_size = 256 \n    n_frames = 10 \n    cnn_features = 256\n    lstm_hidden = 32\n    nfolds = 5 \n    n_epochs = 15\n    TARGET_COLS  = [\n        \"bowel_injury\", \"extravasation_injury\",\n        \"kidney_healthy\", \"kidney_low\", \"kidney_high\",\n        \"liver_healthy\", \"liver_low\", \"liver_high\",\n        \"spleen_healthy\", \"spleen_low\", \"spleen_high\",\n    ]\n    ","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:49:07.7514Z","iopub.execute_input":"2023-09-16T09:49:07.751922Z","iopub.status.idle":"2023-09-16T09:49:07.76563Z","shell.execute_reply.started":"2023-09-16T09:49:07.75189Z","shell.execute_reply":"2023-09-16T09:49:07.764614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class CNN(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.map = nn.Conv2d(in_channels=3, 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","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:49:10.106624Z","iopub.execute_input":"2023-09-16T09:49:10.106971Z","iopub.status.idle":"2023-09-16T09:49:10.114348Z","shell.execute_reply.started":"2023-09-16T09:49:10.106943Z","shell.execute_reply":"2023-09-16T09:49:10.113407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"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        \n        self.fc = nn.Sequential(\n            nn.Linear(CFG.lstm_hidden, 11),\n            nn.Softmax(dim = 1)\n        )\n         \n    def forward(self, x):\n\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":"2023-09-16T09:49:12.236476Z","iopub.execute_input":"2023-09-16T09:49:12.236913Z","iopub.status.idle":"2023-09-16T09:49:12.245976Z","shell.execute_reply.started":"2023-09-16T09:49:12.236882Z","shell.execute_reply":"2023-09-16T09:49:12.244833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''model = Model()\nx = torch.zeros((8, 10, 3, 256, 256))\nt = time.time()\nout = model(x)\nprint(time.time()-t)\nprint(out.shape)'''","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:50:55.086526Z","iopub.execute_input":"2023-09-16T09:50:55.08693Z","iopub.status.idle":"2023-09-16T09:50:55.094136Z","shell.execute_reply.started":"2023-09-16T09:50:55.086899Z","shell.execute_reply":"2023-09-16T09:50:55.092913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Image Handling","metadata":{}},{"cell_type":"code","source":"'''def decode_image(image_path):\n    file_bytes = tf.io.read_file(image_path)\n    image = tf.io.decode_png(file_bytes, channels=3, dtype=tf.uint8)\n    image = tf.image.resize(image, [CFG.img_size,CFG.img_size], method=\"bilinear\")\n    image = tf.cast(image, tf.float32) / 255.0\n    image = tf.transpose(image, perm=[2, 0, 1])\n    return image'''","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:48:05.349072Z","iopub.execute_input":"2023-09-16T09:48:05.349433Z","iopub.status.idle":"2023-09-16T09:48:05.357414Z","shell.execute_reply.started":"2023-09-16T09:48:05.3494Z","shell.execute_reply":"2023-09-16T09:48:05.356497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def decode_image(image_path):\n    image = Image.open(image_path)\n    if image.mode != 'RGB':\n        image = image.convert('RGB')\n    transform = transforms.Compose([\n        transforms.Resize((CFG.img_size, CFG.img_size)),\n        transforms.ToTensor()\n    ])\n    \n    image = transform(image)\n    return image\n","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:01.304456Z","iopub.execute_input":"2023-09-16T09:51:01.305152Z","iopub.status.idle":"2023-09-16T09:51:01.311084Z","shell.execute_reply.started":"2023-09-16T09:51:01.305117Z","shell.execute_reply":"2023-09-16T09:51:01.310039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"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","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:02.861647Z","iopub.execute_input":"2023-09-16T09:51:02.862839Z","iopub.status.idle":"2023-09-16T09:51:02.869632Z","shell.execute_reply.started":"2023-09-16T09:51:02.8628Z","shell.execute_reply":"2023-09-16T09:51:02.868505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def uniform_temporal_subsample(x, num_samples):\n\n    t = len(x)\n    indices = torch.linspace(0, t - 1, num_samples)\n    indices = torch.clamp(indices, 0, t - 1).long()\n    indices = indices.tolist()\n    x = x.tolist()\n    paths = [x[i] for i in indices]\n    return paths","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:03.98633Z","iopub.execute_input":"2023-09-16T09:51:03.986702Z","iopub.status.idle":"2023-09-16T09:51:03.993613Z","shell.execute_reply.started":"2023-09-16T09:51:03.986673Z","shell.execute_reply":"2023-09-16T09:51:03.992464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\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":"2023-09-16T09:51:05.571516Z","iopub.execute_input":"2023-09-16T09:51:05.571937Z","iopub.status.idle":"2023-09-16T09:51:06.544191Z","shell.execute_reply.started":"2023-09-16T09:51:05.571906Z","shell.execute_reply":"2023-09-16T09:51:06.543079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preparing The Data","metadata":{}},{"cell_type":"code","source":"BASE_PATH = f\"/kaggle/input/rsna-atd-512x512-png-v2-dataset\"\ndataframe = pd.read_csv(f\"{BASE_PATH}/train.csv\")\ntrain =  pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/train.csv\")\ndataframe[\"image_path\"] = f\"{BASE_PATH}/train_images\"\\\n                    + \"/\" + dataframe.patient_id.astype(str)\\\n                    + \"/\" + dataframe.series_id.astype(str)\\\n                    + \"/\" + dataframe.instance_number.astype(str) +\".png\"\ndataframe = dataframe.drop_duplicates()\n\ndataframe.head(2)","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:07.609244Z","iopub.execute_input":"2023-09-16T09:51:07.60982Z","iopub.status.idle":"2023-09-16T09:51:07.781361Z","shell.execute_reply.started":"2023-09-16T09:51:07.609782Z","shell.execute_reply":"2023-09-16T09:51:07.780109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataRetriever(Dataset):\n    def __init__(self, paths, transform=None):\n        self.paths = paths\n        self.transform = transform\n          \n    def __len__(self):\n        return len(self.paths)\n    \n    def read_video(self, vid_paths):\n        video = [decode_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].numpy())[\"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,3, CFG.img_size, CFG.img_size)\n        else:\n            video = torch.stack(video) \n        return video    \n    def __getitem__(self, index):\n        _id = self.paths[index]\n        #print(f\"Index: {index}, _id: {_id}\")\n        t_paths = dataframe[dataframe.patient_id == _id].image_path\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,3, CFG.img_size, CFG.img_size)\n        elif channel.shape[0] < CFG.n_frames:\n            pad_frames = CFG.n_frames - channel.shape[0]\n            channel = torch.cat([channel, torch.zeros(pad_frames, 3, CFG.img_size, CFG.img_size)], dim=0)\n            \n        filtered_rows = train[train['patient_id'] == _id]\n        selected_columns = filtered_rows[CFG.TARGET_COLS]\n        y = torch.tensor(selected_columns.values, dtype=torch.float)\n\n        return {\"X\": channel.float(), \"y\": y}\n        ","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:09.835474Z","iopub.execute_input":"2023-09-16T09:51:09.836649Z","iopub.status.idle":"2023-09-16T09:51:09.850202Z","shell.execute_reply.started":"2023-09-16T09:51:09.836605Z","shell.execute_reply":"2023-09-16T09:51:09.849061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''train_retriever = DataRetriever(\n        train[\"patient_id\"].values, \n        train_transform\n    )\ntrain_retriever[0]['X'].shape'''","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:15.485231Z","iopub.execute_input":"2023-09-16T09:51:15.485917Z","iopub.status.idle":"2023-09-16T09:51:15.492789Z","shell.execute_reply.started":"2023-09-16T09:51:15.485878Z","shell.execute_reply":"2023-09-16T09:51:15.49179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''train_retriever = DataRetriever(\n        train[\"patient_id\"].values, \n        train_transform\n    )\nfor idx, dat in enumerate(train_retriever):\n    print('{} {} {}'.format(idx, dat['X'].shape, dat['y']))'''","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:17.465105Z","iopub.execute_input":"2023-09-16T09:51:17.465471Z","iopub.status.idle":"2023-09-16T09:51:17.471764Z","shell.execute_reply.started":"2023-09-16T09:51:17.465441Z","shell.execute_reply":"2023-09-16T09:51:17.470627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Tranning","metadata":{}},{"cell_type":"code","source":"'''class LossMeter:\n    def __init__(self):\n        self.losses = {}\n\n    def update(self, variable_name, loss_value):\n        if variable_name not in self.losses:\n            self.losses[variable_name] = {\"total_loss\": 0, \"total_samples\": 0}\n\n        self.losses[variable_name][\"total_loss\"] += loss_value\n        self.losses[variable_name][\"total_samples\"] += 1\n\n    def get_average_loss(self, variable_name):\n        if variable_name in self.losses:\n            total_loss = self.losses[variable_name][\"total_loss\"]\n            total_samples = self.losses[variable_name][\"total_samples\"]\n            return total_loss / total_samples if total_samples > 0 else 0.0\n        else:\n            return 0.0  # Return 0 if the variable doesn't exist in the dictionary\n\n'''","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:19.701869Z","iopub.execute_input":"2023-09-16T09:51:19.702239Z","iopub.status.idle":"2023-09-16T09:51:19.710014Z","shell.execute_reply.started":"2023-09-16T09:51:19.70221Z","shell.execute_reply":"2023-09-16T09:51:19.70903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LossMeter:\n    def __init__(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","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:21.540819Z","iopub.execute_input":"2023-09-16T09:51:21.541784Z","iopub.status.idle":"2023-09-16T09:51:21.54822Z","shell.execute_reply.started":"2023-09-16T09:51:21.54174Z","shell.execute_reply":"2023-09-16T09:51:21.547085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" '''class AccMeter:\n    def __init__(self):\n        self.accuracies = {}   \n\n    def update(self, variable_name, y_true, y_pred):\n        if variable_name not in self.accuracies:\n            self.accuracies[variable_name] = {\"correct\": 0, \"total\": 0}\n\n        y_true = y_true.cpu().numpy().astype(int)\n        y_pred = np.argmax(y_pred.cpu().numpy(), axis=1)  # Convert softmax outputs to class predictions\n        correct_count = np.sum(y_true == y_pred)\n\n        # Incremental update for the specific variable\n        self.accuracies[variable_name][\"correct\"] += correct_count\n        self.accuracies[variable_name][\"total\"] += len(y_true)\n\n    def get_accuracy(self, variable_name):\n        if variable_name in self.accuracies:\n            return self.accuracies[variable_name][\"correct\"] / self.accuracies[variable_name][\"total\"]\n        else:\n            return 0.0   \n\n'''\n        ","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:23.195723Z","iopub.execute_input":"2023-09-16T09:51:23.196854Z","iopub.status.idle":"2023-09-16T09:51:23.205986Z","shell.execute_reply.started":"2023-09-16T09:51:23.196814Z","shell.execute_reply":"2023-09-16T09:51:23.204775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AccMeter:\n    def __init__(self):\n        self.accuracies = {}\n\n    def update(self, variable_name, y_true, y_pred):\n        if variable_name not in self.accuracies:\n            self.accuracies[variable_name] = {\"correct\": 0, \"total\": 0}\n\n        y_true = y_true.cpu().numpy().astype(int)\n        y_pred = torch.sigmoid(y_pred).cpu().numpy()   \n        predicted_labels = (y_pred > 0.5).astype(int)   \n\n        correct_count = np.sum(y_true == predicted_labels)\n\n        self.accuracies[variable_name][\"correct\"] += correct_count\n        self.accuracies[variable_name][\"total\"] += len(y_true)\n\n    def get_accuracy(self,variable_name):\n        accuracies = {}\n\n        for variable_name, stats in self.accuracies.items():\n            total = stats[\"total\"]\n            correct = stats[\"correct\"]\n            \n            accuracy = correct / total if total > 0 else 0\n            accuracies[variable_name] = accuracy\n\n        return accuracies[variable_name]\n\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Example For Using Loss","metadata":{}},{"cell_type":"code","source":"class Trainer:\n    def __init__(\n        self, \n        model, \n        device, \n        optimizer, \n        criterion, \n        loss_meter, \n        score_meter\n    ):\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 = {'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.n_patience = 0\n        \n        self.messages = {\n            \"epoch\": \"[Epoch {}: {}] loss: {:.5f}, score: {:.5f}, time: {} s\",\n            \"checkpoint\": \"The score improved from {:.5f} to {:.5f}. Save model to '{}'\",\n            \"patience\": \"\\nValid score didn't improve last {} epochs.\"\n        }\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            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            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        train_loss = self.loss_meter()\n        train_scores = self.score_meter()\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            self.optimizer.zero_grad()\n            outputs = self.model(X)\n            \n            loss = self.criterion(outputs, targets.squeeze(dim=1))\n            \n            loss.backward()\n            train_loss.update(loss.detach().item())\n\n            train_scores.update(\"train\",targets.squeeze(dim=1), outputs.detach())\n            \n            self.optimizer.step()\n            _loss = train_loss.avg\n            _score = train_scores.get_accuracy('train')\n        \n            message = 'Train Step {}/{}, train_loss: {:.5f}, train_scores: {}'.format(step, len(train_loader), _loss, _score)\n            self.info_message(message, end=\"\\r\")\n\n            self.optimizer.step()\n        \n        return train_loss.avg, _score, int(time.time() - t)\n    \n    def valid_epoch(self, valid_loader):\n        self.model.eval()\n        t = time.time()\n        valid_loss = self.loss_meter()\n        valid_scores = self.score_meter()\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                outputs = self.model(X) \n                \n                loss = self.criterion(outputs, targets.squeeze(dim=1))\n                \n                valid_loss.update(loss.detach().item())\n                \n                valid_scores.update(\"valid\",targets.squeeze(dim=1), outputs.detach())\n            \n                _score = valid_scores.get_accuracy('valid')\n            message = 'Valid Step {}/{}, Valid_loss: {:.5f}, Valid_scores: {}'.format(step, len(valid_loader), valid_loss.avg, _score)\n            self.info_message(message, end=\"\\r\")\n        \n        return valid_loss.avg,_score, int(time.time() - t)\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\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)","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:25.170799Z","iopub.execute_input":"2023-09-16T09:51:25.171152Z","iopub.status.idle":"2023-09-16T09:51:25.199177Z","shell.execute_reply.started":"2023-09-16T09:51:25.171123Z","shell.execute_reply":"2023-09-16T09:51:25.198191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef split_group(group, test_size=0.2):\n    if len(group) == 1:\n        return (group, pd.DataFrame()) if np.random.rand() < test_size else (pd.DataFrame(), group)\n    else:\n        return train_test_split(group, test_size=test_size, random_state=42)\n    \n\ntrain_data = pd.DataFrame()\nval_data = pd.DataFrame()\nfor _, group in train.groupby(CFG.TARGET_COLS):\n    train_group, val_group= split_group(group)\n    train_data = pd.concat([train_data, train_group], ignore_index = True)\n    val_data = pd.concat([val_data, val_group], ignore_index = True)\n\ntrain_data = train_data[train_data['patient_id'].isin(dataframe['patient_id'])]\nval_data = val_data[val_data['patient_id'].isin(dataframe['patient_id'])]\nstart_time = time.time()\ntrain_retriever = DataRetriever(\n        train_data[\"patient_id\"].values, \n        train_transform\n    )\nval_retriever = DataRetriever(\n        val_data[\"patient_id\"].values\n    )\ntrain_loader = torch_data.DataLoader(\n        train_retriever,\n        batch_size=8,\n        shuffle=True,\n        num_workers=2,\n    )\nvalid_loader = torch_data.DataLoader(\n        val_retriever, \n        batch_size=8,\n        shuffle=False,\n        num_workers=2,\n    )","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:32.311621Z","iopub.execute_input":"2023-09-16T09:51:32.31197Z","iopub.status.idle":"2023-09-16T09:51:32.404547Z","shell.execute_reply.started":"2023-09-16T09:51:32.311943Z","shell.execute_reply":"2023-09-16T09:51:32.403623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''# Iterate through the train_loader and print batch contents\nfor batch in train_loader:\n    # Print the contents of the current batch\n    print(batch[\"X\"].shape)\n'''","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:51:35.053274Z","iopub.execute_input":"2023-09-16T09:51:35.053651Z","iopub.status.idle":"2023-09-16T09:51:35.060202Z","shell.execute_reply.started":"2023-09-16T09:51:35.05362Z","shell.execute_reply":"2023-09-16T09:51:35.059127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Model()\nmodel.to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=0.0001)\ncriterion = nn.BCEWithLogitsLoss()   \n\ntrainer = Trainer(\n        model, \n        device, \n        optimizer, \n        criterion,\n        LossMeter, \n        AccMeter\n    )\nloss, score = trainer.fit(\n        CFG.n_epochs, \n        train_loader, \n        valid_loader, \n        \"best-model.pth\", \n        100,\n    )\ntrainer.plot_loss()\ntrainer.plot_score()\nelapsed_time = time.time() - start_time\nprint('\\nTraining complete in {:.0f}m {:.0f}s'.format(elapsed_time // 60, elapsed_time % 60))\nprint('loss {}'.format(loss))\nprint('score {}'.format(score))","metadata":{"execution":{"iopub.status.busy":"2023-09-16T09:59:40.104763Z","iopub.execute_input":"2023-09-16T09:59:40.105122Z","iopub.status.idle":"2023-09-16T09:59:44.253457Z","shell.execute_reply.started":"2023-09-16T09:59:40.105093Z","shell.execute_reply":"2023-09-16T09:59:44.251524Z"},"trusted":true},"execution_count":null,"outputs":[]}]}