{"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 pydicom as dicom\ndef load_dicom(path,img_size=384):\n    img=dicom.dcmread(path)\n    img.PhotometricInterpretation='YBR_FULL'\n    data=img.pixel_array\n    data=cv2.resize(data,(384,384))\n    return data","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-10-16T06:13:34.010426Z","iopub.execute_input":"2022-10-16T06:13:34.010851Z","iopub.status.idle":"2022-10-16T06:13:34.150172Z","shell.execute_reply.started":"2022-10-16T06:13:34.010748Z","shell.execute_reply":"2022-10-16T06:13:34.149218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -qU ../input/for-pydicom/python_gdcm-3.0.14-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl ../input/for-pydicom/pylibjpeg-1.4.0-py3-none-any.whl --find-links frozen_packages --no-index","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Libraries\nimport os\nimport re\nimport gc\nimport cv2\nimport wandb\nfrom PIL import Image\nimport random\nimport math\nimport shutil\nimport glob\nfrom tqdm import tqdm\nfrom pprint import pprint\nfrom time import time\nimport warnings\nimport pandas as pd\nimport numpy as np\nimport seaborn as sns\nimport matplotlib as mpl\nfrom matplotlib import cm\nimport matplotlib.patches as patches\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nfrom matplotlib.offsetbox import AnnotationBbox, OffsetImage\nfrom matplotlib.colors import ListedColormap, LinearSegmentedColormap\nfrom matplotlib.patches import Rectangle\nfrom IPython.display import display_html\nplt.rcParams.update({'font.size': 16})\n\n# Environment check\nwarnings.filterwarnings(\"ignore\")\nos.environ[\"WANDB_SILENT\"] = \"true\"\nCONFIG = {'competition': 'RSNA_SpineFructure', '_wandb_kernel': 'aot'}\n\n# Custom colors\nclass clr:\n    S = '\\033[1m' + '\\033[94m'\n    E = '\\033[0m'\n    \nmy_colors = [\"#5EAFD9\", \"#449DD1\", \"#3977BB\", \n             \"#2D51A5\", \"#5C4C8F\", \"#8B4679\",\n             \"#C53D4C\", \"#E23836\", \"#FF4633\", \"#FF5746\"]\nCMAP1 = ListedColormap(my_colors)\n\nprint(clr.S+\"Notebook Color Schemes:\"+clr.E)\nsns.palplot(sns.color_palette(my_colors))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:00.840938Z","iopub.execute_input":"2022-10-16T06:14:00.841332Z","iopub.status.idle":"2022-10-16T06:14:01.682149Z","shell.execute_reply.started":"2022-10-16T06:14:00.841284Z","shell.execute_reply":"2022-10-16T06:14:01.680392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -qq einops\n!pip install -qq torchsummary\n\nimport os\nimport cv2\n\nimport torch\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils import shuffle\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\nfrom torch import nn\nfrom torch import Tensor\nfrom torch.utils.data import Subset\n\nfrom PIL import Image\nimport pandas as pd\nimport numpy as np\n\nfrom torchvision.transforms import Compose, Resize, ToTensor\nfrom torch.optim.lr_scheduler import StepLR\nimport torch.optim as optim\n\nfrom einops import rearrange, reduce, repeat\nfrom einops.layers.torch import Rearrange, Reduce","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:01.684677Z","iopub.execute_input":"2022-10-16T06:14:01.685066Z","iopub.status.idle":"2022-10-16T06:14:22.107043Z","shell.execute_reply.started":"2022-10-16T06:14:01.685025Z","shell.execute_reply":"2022-10-16T06:14:22.105928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# PyTorch\nimport torch\nfrom torch.utils.data import TensorDataset, DataLoader, Dataset\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data.sampler import SubsetRandomSampler, RandomSampler, SequentialSampler\nfrom torch.optim.lr_scheduler import StepLR, ReduceLROnPlateau, CosineAnnealingLR\nimport torchvision\nimport torchvision.transforms as transforms\nimport albumentations\n\nfrom sklearn.model_selection import GroupKFold, train_test_split, StratifiedKFold\nfrom sklearn.metrics import roc_auc_score, cohen_kappa_score, confusion_matrix\n","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:22.111045Z","iopub.execute_input":"2022-10-16T06:14:22.111599Z","iopub.status.idle":"2022-10-16T06:14:22.120491Z","shell.execute_reply.started":"2022-10-16T06:14:22.111564Z","shell.execute_reply":"2022-10-16T06:14:22.119093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport numpy as np\n\ndef load_dicom_3d(patient_id,num_imgs=96,img_size=384):\n    files=sorted(glob.glob(f\"{data_dir}/{patient_id}/*.dcm\"))\n    middle=len(files)//2\n    num_imgs2=num_imgs//2\n    p1=max(0,middle-num_imgs2)\n    p2=min(len(files),middle+num_imgs2)\n    img3d=np.stack([load_dicom(f) for f in files[p1:p2]]).T\n    if img3d.shape[-1]<num_imgs:\n        n_zero=np.zeros((img_size,img_size,num_imgs-img3d.shape[-1]))\n        img3d=np.concatenate((img3d,n_zero),axis=-1)\n    \n    if np.min(img3d)<np.max(img3d):\n        img3d=img3d-np.min(img3d)\n        img3d=img3d/np.max(img3d)\n        \n    return np.expand_dims(img3d,0)\n","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:22.123544Z","iopub.execute_input":"2022-10-16T06:14:22.123932Z","iopub.status.idle":"2022-10-16T06:14:22.137721Z","shell.execute_reply.started":"2022-10-16T06:14:22.123894Z","shell.execute_reply":"2022-10-16T06:14:22.136582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir='../input/rsna-2022-cervical-spine-fracture-detection/train_images'","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:22.139045Z","iopub.execute_input":"2022-10-16T06:14:22.139705Z","iopub.status.idle":"2022-10-16T06:14:22.149411Z","shell.execute_reply.started":"2022-10-16T06:14:22.139669Z","shell.execute_reply":"2022-10-16T06:14:22.148486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data=load_dicom_3d(\"1.2.826.0.1.3680043.10001\")\ndata.shape","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:22.151061Z","iopub.execute_input":"2022-10-16T06:14:22.151464Z","iopub.status.idle":"2022-10-16T06:14:22.486615Z","shell.execute_reply.started":"2022-10-16T06:14:22.151428Z","shell.execute_reply":"2022-10-16T06:14:22.485485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install '../input/einops/einops-0.3.0-py2.py3-none-any.whl'\nfrom einops import rearrange, repeat\nfrom einops.layers.torch import Rearrange\n\n# helpers\n\ndef pair(t):\n    return t if isinstance(t, tuple) else (t, t)\n\n# classes\n\nclass PreNorm(nn.Module):\n    def __init__(self, dim, fn):\n        super().__init__()\n        self.norm = nn.LayerNorm(dim)\n        self.fn = fn\n    def forward(self, x, **kwargs):\n        return self.fn(self.norm(x), **kwargs)\n\nclass FeedForward(nn.Module):\n    def __init__(self, dim, hidden_dim, dropout = 0.):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(dim, hidden_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, dim),\n            nn.Dropout(dropout)\n        )\n    def forward(self, x):\n        return self.net(x)\n\nclass Attention(nn.Module):\n    def __init__(self, dim, heads = 8, dim_head = 64, dropout = 0.):\n        super().__init__()\n        inner_dim = dim_head *  heads\n        project_out = not (heads == 1 and dim_head == dim)\n\n        self.heads = heads\n        self.scale = dim_head ** -0.5\n\n        self.attend = nn.Softmax(dim = -1)\n        self.to_qkv = nn.Linear(dim, inner_dim * 3, bias = False)\n\n        self.to_out = nn.Sequential(\n            nn.Linear(inner_dim, dim),\n            nn.Dropout(dropout)\n        ) if project_out else nn.Identity()\n\n    def forward(self, x):\n        qkv = self.to_qkv(x).chunk(3, dim = -1)\n        q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = self.heads), qkv)\n\n        dots = torch.matmul(q, k.transpose(-1, -2)) * self.scale\n\n        attn = self.attend(dots)\n\n        out = torch.matmul(attn, v)\n        out = rearrange(out, 'b h n d -> b n (h d)')\n        return self.to_out(out)\n\nclass Transformer(nn.Module):\n    def __init__(self, dim, depth, heads, dim_head, mlp_dim, dropout = 0.):\n        super().__init__()\n        self.layers = nn.ModuleList([])\n        mlp_dim = 2048\n        for _ in range(depth):\n            #print (dim, mlp_dim)\n            self.layers.append(nn.ModuleList([\n                PreNorm(dim, Attention(dim, heads = heads, dim_head = dim_head, dropout = dropout)),\n                PreNorm(dim, FeedForward(dim, mlp_dim, dropout = dropout))\n            ]))\n    def forward(self, x):\n        for attn, ff in self.layers:\n            x = attn(x) + x\n            x = ff(x) + x\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:22.48846Z","iopub.execute_input":"2022-10-16T06:14:22.488857Z","iopub.status.idle":"2022-10-16T06:14:31.94576Z","shell.execute_reply.started":"2022-10-16T06:14:22.488808Z","shell.execute_reply":"2022-10-16T06:14:31.944616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, *, image_size, patch_size, num_classes, dim, depth, heads, mlp_dim, channels = 3, dropout = 0., emb_dropout = 0.):\n        super().__init__()\n        assert image_size % patch_size == 0, 'image dimensions must be divisible by the patch size'\n        num_patches = (image_size // patch_size) *(image_size // patch_size)* 2\n        patch_dim = channels * patch_size ** 3\n\n        self.patch_size = patch_size\n\n        self.pos_embedding = nn.Parameter(torch.randn(1, num_patches + 1, dim))\n        self.patch_to_embedding = nn.Linear(patch_dim, dim)\n        self.cls_token = nn.Parameter(torch.randn(1, 1, dim))\n        self.dropout = nn.Dropout(emb_dropout)\n        #print (mlp_dim)\n        self.transformer = Transformer(dim, depth, heads, mlp_dim, dropout)\n        #print (dim)\n        self.to_cls_token = nn.Identity()\n\n        self.mlp_head = nn.Sequential(\n            nn.LayerNorm(dim),\n            nn.Linear(dim, mlp_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(mlp_dim, num_classes),\n            nn.Dropout(dropout),\n            nn.Softmax(dim=1)\n            # add the line for sigmoid layer \n        )\n\n    def forward(self, img, mask = None):\n        p = self.patch_size\n        #print (img.shape)\n        x = rearrange(img, 'b c (h p1) (w p2) (d p3) -> b (h w d) (p1 p2 p3 c)', p1 = p, p2 = p, p3 = p)\n        #print (x.shape)\n        x = self.patch_to_embedding(x)\n        #print (x.shape)\n        cls_tokens = self.cls_token.expand(img.shape[0], -1, -1)\n        #print (cls_tokens.shape)\n        x = torch.cat((cls_tokens, x), dim=1)\n        #print (x.shape)\n        #print (self.pos_embedding.shape)\n        x += self.pos_embedding\n        x = self.dropout(x)\n\n        x = self.transformer(x)\n\n        x = self.to_cls_token(x[:, 0])\n        return self.mlp_head(x)","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:31.947605Z","iopub.execute_input":"2022-10-16T06:14:31.950656Z","iopub.status.idle":"2022-10-16T06:14:31.96237Z","shell.execute_reply.started":"2022-10-16T06:14:31.95062Z","shell.execute_reply":"2022-10-16T06:14:31.961096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:31.964487Z","iopub.execute_input":"2022-10-16T06:14:31.964963Z","iopub.status.idle":"2022-10-16T06:14:32.002204Z","shell.execute_reply.started":"2022-10-16T06:14:31.964924Z","shell.execute_reply":"2022-10-16T06:14:32.001195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_cols = ['C1', 'C2', 'C3',\n               'C4', 'C5', 'C6', 'C7','patient_overall']","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:32.006867Z","iopub.execute_input":"2022-10-16T06:14:32.008032Z","iopub.status.idle":"2022-10-16T06:14:32.013586Z","shell.execute_reply.started":"2022-10-16T06:14:32.007994Z","shell.execute_reply":"2022-10-16T06:14:32.012603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install monai","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:32.016647Z","iopub.execute_input":"2022-10-16T06:14:32.016987Z","iopub.status.idle":"2022-10-16T06:14:41.88047Z","shell.execute_reply.started":"2022-10-16T06:14:32.01694Z","shell.execute_reply":"2022-10-16T06:14:41.879291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import TensorDataset, DataLoader, Dataset\nfrom monai.transforms import Randomizable, apply_transform\nnatsort = lambda s: [int(t) if t.isdigit() else t.lower() for t in re.split('(\\d+)', s)]\n\nclass RSNADataset(Dataset, Randomizable):\n    \n    def __init__(self, csv, mode, transform=None):\n        self.csv = csv\n        self.mode = mode\n        self.transform = transform\n        \n    def __len__(self):\n        return self.csv.shape[0]\n    \n    def randomize(self) -> None:\n        '''-> None is a type annotation for the function that states \n        that this function returns None.'''\n        \n        MAX_SEED = np.iinfo(np.uint32).max + 1\n        self.seed = self.R.randint(MAX_SEED, dtype=\"uint32\")\n        \n    def __getitem__(self, index):\n        # Set Random Seed\n        self.randomize()\n        \n        dt = self.csv.iloc[index, :]\n        study_paths = glob.glob(f\"train_DICOM/{dt.StudyInstanceUID}/*\")\n        study_paths.sort(key=natsort)\n        \n        # Load images\n        stacked_image=load_dicom_3d(dt.StudyInstanceUID)\n        #print(\"need to sqz shape\",stacked_image.shape)\n        \n        if self.transform:\n            if isinstance(self.transform, Randomizable):\n                self.transform.set_random_state(seed=self.seed)\n                \n            stacked_image = apply_transform(self.transform, stacked_image)\n        \n        if self.mode==\"test\":\n            return {\"X\":torch.tensor(stacked_image).float(),\"id\":dt.StudyInstanceUID}\n        else:\n            targets = torch.tensor(dt[target_cols]).float()\n            return {\"X\": torch.tensor(stacked_image).float(),\"y\":targets}\n        ","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:41.882587Z","iopub.execute_input":"2022-10-16T06:14:41.883014Z","iopub.status.idle":"2022-10-16T06:14:43.135644Z","shell.execute_reply.started":"2022-10-16T06:14:41.882968Z","shell.execute_reply":"2022-10-16T06:14:43.134469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Local Setup\nDF_SIZE = 0.03\nN_SPLITS = 5\nKERNEL_TYPE = 'ViT3d'\nIMG_RESIZE = 384\nSTACK_RESIZE = 50\nuse_amp = False\nNUM_WORKERS = 4\nBATCH_SIZE = 8\nLR = 0.0005\nOUT_DIM = 8\nEPOCHS = 10","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:43.137627Z","iopub.execute_input":"2022-10-16T06:14:43.138164Z","iopub.status.idle":"2022-10-16T06:14:43.144114Z","shell.execute_reply.started":"2022-10-16T06:14:43.138119Z","shell.execute_reply":"2022-10-16T06:14:43.142708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.random.seed(0)\nimport pandas as pd\nimport random\nfrom sklearn.model_selection import GroupKFold, train_test_split, StratifiedKFold\n\ndf = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n\n# Sample down df\ninstances = df.StudyInstanceUID.unique().tolist()\ninstances = random.sample(instances, k=int(len(instances)*DF_SIZE))\ndf = df[df[\"StudyInstanceUID\"].isin(instances)].reset_index(drop=True)\nprint(\"Dataframe size:\", df.shape)\n\n# Create folds\nkfold = GroupKFold(n_splits=N_SPLITS)\ndf['fold'] = -1\n\n# Append fold\nfor k, (_, valid_i) in enumerate(kfold.split(df,\n                                             groups=df.StudyInstanceUID)):\n    df.loc[valid_i, 'fold'] = k\n    \nprint(\"K Folds Count:\")\ndf[\"fold\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:43.145847Z","iopub.execute_input":"2022-10-16T06:14:43.146608Z","iopub.status.idle":"2022-10-16T06:14:43.177477Z","shell.execute_reply.started":"2022-10-16T06:14:43.14657Z","shell.execute_reply":"2022-10-16T06:14:43.176472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def data_to_device(data):\n    X, y = data.values()\n    return X.to(device), y.to(device)","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:43.178982Z","iopub.execute_input":"2022-10-16T06:14:43.179339Z","iopub.status.idle":"2022-10-16T06:14:43.185204Z","shell.execute_reply.started":"2022-10-16T06:14:43.179295Z","shell.execute_reply":"2022-10-16T06:14:43.184086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Sample data\nsample_df = df.head(6)\n\n# Instantiate Dataset object\ndataset = RSNADataset(csv=sample_df, mode=\"train\")\n\n# The Dataloader\ndataloader = DataLoader(dataset, batch_size=3, shuffle=False)\n\n# Output of the Dataloader\nfor k, data in enumerate(dataloader):\n    image, targets = data_to_device(data)\n    #img_u=torch.mean(img,0)\n    #img_u=img_u.permute(0,1,2)\n    print(clr.S + f\"Batch: {k}\" + clr.E, \"\\n\" +\n          clr.S + \"Image:\" + clr.E, image.shape, \"\\n\" +\n          clr.S + \"Targets:\" + clr.E, targets, \"\\n\" +\n          \"=\"*50)","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:43.186697Z","iopub.execute_input":"2022-10-16T06:14:43.187236Z","iopub.status.idle":"2022-10-16T06:14:53.313249Z","shell.execute_reply.started":"2022-10-16T06:14:43.187197Z","shell.execute_reply":"2022-10-16T06:14:53.312045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:53.314842Z","iopub.execute_input":"2022-10-16T06:14:53.315535Z","iopub.status.idle":"2022-10-16T06:14:53.540182Z","shell.execute_reply.started":"2022-10-16T06:14:53.315495Z","shell.execute_reply":"2022-10-16T06:14:53.538954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def add_in_file(text, f):\n    \n    with open(f'log_{KERNEL_TYPE}.txt', 'a+') as f:\n        print(text, file=f)","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:53.541766Z","iopub.execute_input":"2022-10-16T06:14:53.542385Z","iopub.status.idle":"2022-10-16T06:14:53.54985Z","shell.execute_reply.started":"2022-10-16T06:14:53.542346Z","shell.execute_reply":"2022-10-16T06:14:53.54895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.nn import functional as torch_functional\ncriterion = torch_functional.binary_cross_entropy_with_logits","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:53.551262Z","iopub.execute_input":"2022-10-16T06:14:53.551763Z","iopub.status.idle":"2022-10-16T06:14:53.560549Z","shell.execute_reply.started":"2022-10-16T06:14:53.551728Z","shell.execute_reply":"2022-10-16T06:14:53.559532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_epoch(model, dataloader,optimizer,epochs, f,fold):\n    \n    # Add info to file\n    print(\"Training...\")\n    add_in_file('Training...', f)\n    \n    # Track training time for 1 epoch\n    start_time = time()\n    train = df[df[\"fold\"] != fold].reset_index(drop=True)\n    train_dataset = RSNADataset(csv=train, mode=\"train\")\n    trainloader = DataLoader(train_dataset, batch_size=BATCH_SIZE,\n                             sampler=RandomSampler(train_dataset))\n    # === TRAIN ===\n    model.train()\n    \n    train_losses, train_comp_losses = [], []\n    for epoch in range(epochs):\n        for i, data, in enumerate(trainloader):\n            image, targets = data_to_device(data)\n            output = model(image)\n            #targets=torch.argmax(targets,dim=1)\n            loss = criterion(output,targets)\n            optimizer.zero_grad()\n            loss.sum().backward()\n            optimizer.step()\n            train_losses.append(loss.detach().cpu().numpy())\n            gc.collect()\n    try:\n        mean_train_loss = np.mean(train_losses)\n    except ValueError:\n        mean_train_loss=0.5\n        print(\"mean_train_loss is not able to found here.\")\n        \n   # print(\"train mean losses shape\",mean_train_loss.shape)\n    print(\"train_loss\",mean_train_loss)\n    gc.collect()\n    return mean_train_loss","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:53.562101Z","iopub.execute_input":"2022-10-16T06:14:53.56283Z","iopub.status.idle":"2022-10-16T06:14:53.573515Z","shell.execute_reply.started":"2022-10-16T06:14:53.562776Z","shell.execute_reply":"2022-10-16T06:14:53.57252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def valid_epoch(model, dataloader, epoch, f,fold):\n    \n    # Add info to file\n    print(\"Validation...\")\n    add_in_file('Validation...', f)\n    \n    # Track validation time for 1 epoch\n    start_time = time()\n    \n    # === EVAL ===\n    model.eval()\n    valid = df[df[\"fold\"] == fold].reset_index(drop=True)\n    valid_preds, valid_targets, valid_comp_loss = [], [], []\n    valid_dataset = RSNADataset(csv=valid, mode=\"train\")\n    \n    validloader = DataLoader(valid_dataset, batch_size=BATCH_SIZE)\n    with torch.no_grad():\n\n        for i,data in enumerate(validloader):\n            image, targets = data_to_device(data)\n            #img_u=torch.mean(img,0)\n            #img_u=img_u.permute(0,1,2,3)\n            #img_u = img_u.to(device)\n            #label = label.to(device)\n            logits = model(image)\n            #print(\"valid targets shape\",targets.shape)\n            #targets=torch.argmax(targets,dim=1)\n            #print(\"output\",output)\n            #print(\"targets\",targets)\n            valid_targets.append(targets.detach().cpu())\n            valid_preds.append(logits.detach().cpu())\n             # Overall Valid Loss\n    print(\"valid_preds_shape\",torch.cat(valid_preds).shape)\n    print(\"valid_targets_shape\",torch.cat(valid_targets).shape)\n    valid_losses = criterion(torch.cat(valid_preds), torch.cat(valid_targets)).numpy()\n    \n    try:\n        mean_valid_loss = np.mean(valid_losses)\n    except ValueError:\n        mean_valid_loss=0.5\n        print(\"mean_valid_loss is not able to found here.\")\n    #print(\"mean_valid_loss_shape\",mean_valid_loss.shape)\n    PREDS = np.concatenate(torch.cat(valid_preds).numpy())\n    TARGETS = np.concatenate(torch.cat(valid_targets).numpy())\n    try:\n        print(roc_auc_score(TARGETS, PREDS))\n    except ValueError:\n        pass\n    print(\"valid loss\",mean_valid_loss)\n    gc.collect()\n    return mean_valid_loss","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:53.575126Z","iopub.execute_input":"2022-10-16T06:14:53.575577Z","iopub.status.idle":"2022-10-16T06:14:53.589183Z","shell.execute_reply.started":"2022-10-16T06:14:53.575543Z","shell.execute_reply":"2022-10-16T06:14:53.588045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_train(fold):\n    train = df[df[\"fold\"] != fold].reset_index(drop=True)\n    valid = df[df[\"fold\"] == fold].reset_index(drop=True)\n    train_dataset = RSNADataset(csv=train, mode=\"train\")\n    valid_dataset = RSNADataset(csv=valid, mode=\"train\"\n                            )\n    trainloader = DataLoader(train_dataset, batch_size=BATCH_SIZE,\n                             sampler=RandomSampler(train_dataset))\n    validloader = DataLoader(valid_dataset, batch_size=BATCH_SIZE)\n    \n    model = Model(\n        image_size = 384,\n        patch_size = 48,\n        num_classes = 8,\n        dim = 768,\n        depth = 2,\n        heads = 16,\n        mlp_dim = 1536,\n        channels = 1,\n        dropout = 0.1,\n        emb_dropout = 0.1\n    )\n    model.to(device)\n    optimizer = optim.Adam(model.parameters(), lr=3e-5)\n    scheduler = StepLR(optimizer, step_size=1, gamma=0.7)\n     # Initiate initial loss\n    valid_loss_BEST = 0.71\n    # Create model name\n    model_file = f'best_fold_vit3d_{fold}.pth'\n    # Create file to save outputs\n    f = open(f'log_.txt', 'a')\n    for epoch in range(EPOCHS):\n        add_in_file('======== Epoch: {}/{} ========'.format(epoch+1, EPOCHS), f)\n        print(\"=\"*8, clr.S+f\"Epoch {epoch}\"+clr.E, \"=\"*8)\n               \n        # Train & Validate\n        mean_train_loss = train_epoch(model, trainloader, optimizer, epoch, f,fold)\n        mean_valid_loss = valid_epoch(model, validloader, epoch, f,fold)\n        \n        # Save model\n        if mean_valid_loss <= valid_loss_BEST:\n            print('Saving model ...')\n            add_in_file('Saving model => {}'.format(model_file), f)\n            torch.save(model.state_dict(), model_file)\n            valid_loss_BEST = mean_valid_loss\n            \n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:53.590439Z","iopub.execute_input":"2022-10-16T06:14:53.591155Z","iopub.status.idle":"2022-10-16T06:14:53.605659Z","shell.execute_reply.started":"2022-10-16T06:14:53.591117Z","shell.execute_reply":"2022-10-16T06:14:53.604576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run_train(fold=0)","metadata":{"execution":{"iopub.status.busy":"2022-10-16T06:14:53.607359Z","iopub.execute_input":"2022-10-16T06:14:53.607701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}