{"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":"# MONAI 3D model\n!pip install -q monai\n!pip install -q git+https://github.com/ildoonet/pytorch-gradual-warmup-lr.git","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:10.600827Z","iopub.execute_input":"2022-09-15T08:55:10.601937Z","iopub.status.idle":"2022-09-15T08:55:37.30375Z","shell.execute_reply.started":"2022-09-15T08:55:10.601809Z","shell.execute_reply":"2022-09-15T08:55:37.30253Z"},"trusted":true},"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\nfrom glob import 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-09-15T08:55:37.306359Z","iopub.execute_input":"2022-09-15T08:55:37.306767Z","iopub.status.idle":"2022-09-15T08:55:39.021976Z","shell.execute_reply.started":"2022-09-15T08:55:37.306725Z","shell.execute_reply":"2022-09-15T08:55:39.020662Z"},"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\nfrom warmup_scheduler import GradualWarmupScheduler\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\n# MONAI 3D\nfrom monai.transforms import Randomizable, apply_transform\nfrom monai.transforms import Compose, Resize, ScaleIntensity, ToTensor, RandAffine","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:39.023639Z","iopub.execute_input":"2022-09-15T08:55:39.024341Z","iopub.status.idle":"2022-09-15T08:55:44.552738Z","shell.execute_reply.started":"2022-09-15T08:55:39.024287Z","shell.execute_reply":"2022-09-15T08:55:44.551759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"import wandb\nfrom wandb.keras import WandbCallback\n\nwandb.login()","metadata":{}},{"cell_type":"code","source":"DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(clr.S+\"Device:\"+clr.E, DEVICE)\n\n# Kaggle Notebook Setup\nDF_SIZE = 0.003\nN_SPLITS = 5\nKERNEL_TYPE = 'VIT'\nIMG_RESIZE = 100\nSTACK_RESIZE = 50\nuse_amp = False\nNUM_WORKERS = 1\nBATCH_SIZE = 2\nLR = 0.05\nOUT_DIM = 8\nEPOCHS = 2","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:44.555454Z","iopub.execute_input":"2022-09-15T08:55:44.55641Z","iopub.status.idle":"2022-09-15T08:55:44.627101Z","shell.execute_reply.started":"2022-09-15T08:55:44.556367Z","shell.execute_reply":"2022-09-15T08:55:44.626053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_cols = ['C1', 'C2', 'C3', \n               'C4', 'C5', 'C6', 'C7',\n               'patient_overall']","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:44.628776Z","iopub.execute_input":"2022-09-15T08:55:44.629503Z","iopub.status.idle":"2022-09-15T08:55:44.637903Z","shell.execute_reply.started":"2022-09-15T08:55:44.629458Z","shell.execute_reply":"2022-09-15T08:55:44.636851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# src: https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/341854\ncompetition_weights = {\n    '-' : torch.tensor([1, 1, 1, 1, 1, 1, 1, 7], dtype=torch.float, device=DEVICE),\n    '+' : torch.tensor([2, 2, 2, 2, 2, 2, 2, 14], dtype=torch.float, device=DEVICE),\n}","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:44.639577Z","iopub.execute_input":"2022-09-15T08:55:44.639988Z","iopub.status.idle":"2022-09-15T08:55:47.694618Z","shell.execute_reply.started":"2022-09-15T08:55:44.639951Z","shell.execute_reply":"2022-09-15T08:55:47.693634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Example\n\n# Prediction (very bad)\nlogits = torch.tensor([[0.2221, 0.1037, 0.0739, 0.1112, 0.1026, 0.0902, 0.1597, 0.1365],\n                       [0.1702, 0.0952, 0.0815, 0.1262, 0.1185, 0.1097, 0.1675, 0.1312]],\n                      device=DEVICE)\nprint(clr.S+\"Prediction:\"+clr.E, \"\\n\", logits)\n\n# Actual\ntargets = torch.tensor([[0., 0., 0., 0., 0., 0., 0., 0.],\n                        [1., 0., 0., 0., 0., 0., 0., 1.]], device=DEVICE)\nprint(clr.S+\"Target:\"+clr.E, \"\\n\", targets)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.696174Z","iopub.execute_input":"2022-09-15T08:55:47.696655Z","iopub.status.idle":"2022-09-15T08:55:47.738715Z","shell.execute_reply.started":"2022-09-15T08:55:47.696621Z","shell.execute_reply":"2022-09-15T08:55:47.737661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compute the weights\nweights = targets * competition_weights['+'] + (1 - targets) * competition_weights['-']\nprint(clr.S+\"Weights:\"+clr.E, \"\\n\", weights)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.740256Z","iopub.execute_input":"2022-09-15T08:55:47.740658Z","iopub.status.idle":"2022-09-15T08:55:47.75156Z","shell.execute_reply.started":"2022-09-15T08:55:47.740622Z","shell.execute_reply":"2022-09-15T08:55:47.750542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compute losses on label and exam level\nL = torch.zeros(targets.shape, device=DEVICE)\n\nw = weights\ny = targets\np = logits\n\nfor i in range(L.shape[0]):\n    for j in range(L.shape[1]):\n        L[i, j] = -w[i, j] * (\n            y[i, j] * math.log(p[i, j]) +\n            (1 - y[i, j]) * math.log(1 - p[i, j]))\n        \nprint(clr.S+\"LOSSES:\"+clr.E, \"\\n\", L)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.753378Z","iopub.execute_input":"2022-09-15T08:55:47.754408Z","iopub.status.idle":"2022-09-15T08:55:47.770983Z","shell.execute_reply.started":"2022-09-15T08:55:47.754363Z","shell.execute_reply":"2022-09-15T08:55:47.769779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Average Loss on Exam (or patient)\nExams_Loss = torch.div(torch.sum(L, dim=1), torch.sum(w, dim=1))\n\nprint(clr.S+\"Exam Losses:\"+clr.E, \"\\n\", Exams_Loss)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.774669Z","iopub.execute_input":"2022-09-15T08:55:47.776549Z","iopub.status.idle":"2022-09-15T08:55:47.784194Z","shell.execute_reply.started":"2022-09-15T08:55:47.776522Z","shell.execute_reply":"2022-09-15T08:55:47.783172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_custom_loss(logits, targets):\n    \n    # Compute the weights\n    weights = targets * competition_weights['+'] + (1 - targets) * competition_weights['-']\n    \n    # Losses on label and exam level\n    L = torch.zeros(targets.shape, device=DEVICE)\n\n    w = weights\n    y = targets\n    p = logits\n    eps=1e-8\n\n    for i in range(L.shape[0]):\n        for j in range(L.shape[1]):\n            L[i, j] = -w[i, j] * (\n                y[i, j] * math.log(p[i, j] + eps) +\n                (1 - y[i, j]) * math.log(1 - p[i, j] + eps))\n            \n    # Average Loss on Exam (or patient)\n    Exams_Loss = torch.div(torch.sum(L, dim=1), torch.sum(w, dim=1))\n    \n    return Exams_Loss","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.785845Z","iopub.execute_input":"2022-09-15T08:55:47.786206Z","iopub.status.idle":"2022-09-15T08:55:47.794768Z","shell.execute_reply.started":"2022-09-15T08:55:47.786172Z","shell.execute_reply":"2022-09-15T08:55:47.793712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.random.seed(0)\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(clr.S+\"Dataframe size:\"+clr.E, 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(clr.S+\"K Folds Count:\"+clr.E)\ndf[\"fold\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.796294Z","iopub.execute_input":"2022-09-15T08:55:47.796701Z","iopub.status.idle":"2022-09-15T08:55:47.842218Z","shell.execute_reply.started":"2022-09-15T08:55:47.796666Z","shell.execute_reply":"2022-09-15T08:55:47.841177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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(f\"../input/rsna-fracture-detection/zip_png_images/{dt.StudyInstanceUID}/*\")\n        study_paths.sort()\n        \n        # Load images\n        study_images = [cv2.imread(path)[:,:,::-1] for path in study_paths]\n        # Stack all scans into 1\n        stacked_image = np.stack([img.astype(np.float32) for img in study_images],axis=2).transpose(3,1,0,2)\n        \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 {\"image\": stacked_image}\n        else:\n            targets = torch.tensor(dt[target_cols]).float()\n            return {\"image\": stacked_image,\n                    \"targets\": targets}","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.843971Z","iopub.execute_input":"2022-09-15T08:55:47.84466Z","iopub.status.idle":"2022-09-15T08:55:47.855521Z","shell.execute_reply.started":"2022-09-15T08:55:47.844623Z","shell.execute_reply":"2022-09-15T08:55:47.854478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def data_to_device(data):\n    \n    image, targets = data.values()\n    return image.to(DEVICE), targets.to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.857094Z","iopub.execute_input":"2022-09-15T08:55:47.857786Z","iopub.status.idle":"2022-09-15T08:55:47.86949Z","shell.execute_reply.started":"2022-09-15T08:55:47.857739Z","shell.execute_reply":"2022-09-15T08:55:47.868559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = Compose([ScaleIntensity(),\n                            Resize((IMG_RESIZE, IMG_RESIZE,STACK_RESIZE)),ToTensor()])\nvalid_transforms = Compose([ScaleIntensity(),Resize((IMG_RESIZE, IMG_RESIZE,STACK_RESIZE)),ToTensor()])","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.871081Z","iopub.execute_input":"2022-09-15T08:55:47.872309Z","iopub.status.idle":"2022-09-15T08:55:47.882533Z","shell.execute_reply.started":"2022-09-15T08:55:47.87227Z","shell.execute_reply":"2022-09-15T08:55:47.881545Z"},"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\", transform=train_transforms)\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=torch.mean(image, -1)\n    img=img.permute(0,1,2,3)\n    print(clr.S + f\"Batch: {k}\" + clr.E, \"\\n\" +\n          clr.S + \"Image:\" + clr.E, img.shape, \"\\n\" +\n          clr.S + \"Targets:\" + clr.E, targets, \"\\n\" +\n          \"=\"*50)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:55:47.885967Z","iopub.execute_input":"2022-09-15T08:55:47.886272Z","iopub.status.idle":"2022-09-15T08:56:39.452764Z","shell.execute_reply.started":"2022-09-15T08:55:47.886212Z","shell.execute_reply":"2022-09-15T08:56:39.451691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del dataset, dataloader, image, targets\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:56:39.454182Z","iopub.execute_input":"2022-09-15T08:56:39.455138Z","iopub.status.idle":"2022-09-15T08:56:39.654599Z","shell.execute_reply.started":"2022-09-15T08:56:39.4551Z","shell.execute_reply":"2022-09-15T08:56:39.653351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam\nfrom torch.nn import CrossEntropyLoss\nfrom torch.utils.data import DataLoader","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:56:39.656113Z","iopub.execute_input":"2022-09-15T08:56:39.656563Z","iopub.status.idle":"2022-09-15T08:56:39.663972Z","shell.execute_reply.started":"2022-09-15T08:56:39.656522Z","shell.execute_reply":"2022-09-15T08:56:39.662867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision.transforms import ToTensor\nfrom torchvision.datasets.mnist import MNIST\n\nnp.random.seed(0)\ntorch.manual_seed(0)\n\n\ndef patchify(images, n_patches):\n    n, c, h, w = images.shape\n    print(\"images shape\",images.shape)\n    assert h == w, \"Patchify method is implemented for square images only\"\n\n    patches = torch.zeros(n, n_patches ** 2, h * w // n_patches ** 2)\n    patch_size = h // n_patches\n\n    for idx, image in enumerate(images):\n        for i in range(n_patches):\n            for j in range(n_patches):\n                patch = image[:, i * patch_size: (i + 1) * patch_size, j * patch_size: (j + 1) * patch_size]\n                patches[idx, i * n_patches + j] = patch.flatten().squeeze()\n    return patches\n\n\nclass MyMSA(nn.Module):\n    def __init__(self, d, n_heads=2):\n        super(MyMSA, self).__init__()\n        self.d = d\n        self.n_heads = n_heads\n\n        assert d % n_heads == 0, f\"Can't divide dimension {d} into {n_heads} heads\"\n\n        d_head = int(d / n_heads)\n        self.q_mappings = nn.ModuleList([nn.Linear(d_head, d_head) for _ in range(self.n_heads)])\n        self.k_mappings = nn.ModuleList([nn.Linear(d_head, d_head) for _ in range(self.n_heads)])\n        self.v_mappings = nn.ModuleList([nn.Linear(d_head, d_head) for _ in range(self.n_heads)])\n        self.d_head = d_head\n        self.softmax = nn.Softmax(dim=-1)\n\n    def forward(self, sequences):\n        # Sequences has shape (N, seq_length, token_dim)\n        # We go into shape    (N, seq_length, n_heads, token_dim / n_heads)\n        # And come back to    (N, seq_length, item_dim)  (through concatenation)\n        result = []\n        for sequence in sequences:\n            seq_result = []\n            for head in range(self.n_heads):\n                q_mapping = self.q_mappings[head]\n                k_mapping = self.k_mappings[head]\n                v_mapping = self.v_mappings[head]\n\n                seq = sequence[:, head * self.d_head: (head + 1) * self.d_head]\n                q, k, v = q_mapping(seq), k_mapping(seq), v_mapping(seq)\n\n                attention = self.softmax(q @ k.T / (self.d_head ** 0.5))\n                seq_result.append(attention @ v)\n            result.append(torch.hstack(seq_result))\n        return torch.cat([torch.unsqueeze(r, dim=0) for r in result])\n\n\nclass MyViTBlock(nn.Module):\n    def __init__(self, hidden_d, n_heads, mlp_ratio=4):\n        super(MyViTBlock, self).__init__()\n        self.hidden_d = hidden_d\n        self.n_heads = n_heads\n\n        self.norm1 = nn.LayerNorm(hidden_d)\n        self.mhsa = MyMSA(hidden_d, n_heads)\n        self.norm2 = nn.LayerNorm(hidden_d)\n        self.mlp = nn.Sequential(\n            nn.Linear(hidden_d, mlp_ratio * hidden_d),\n            nn.GELU(),\n            nn.Linear(mlp_ratio * hidden_d, hidden_d)\n        )\n\n    def forward(self, x):\n        out = x + self.mhsa(self.norm1(x))\n        out = out + self.mlp(self.norm2(out))\n        return out\n\n\nclass MyViT(nn.Module):\n    def __init__(self, chw, n_patches=7, n_blocks=2, hidden_d=8, n_heads=2, out_d=10):\n        # Super constructor\n        super(MyViT, self).__init__()\n        \n        # Attributes\n        self.chw = chw # ( C , H , W )\n        self.n_patches = n_patches\n        self.n_blocks = n_blocks\n        self.n_heads = n_heads\n        self.hidden_d = hidden_d\n        \n        # Input and patches sizes\n        assert chw[1] % n_patches == 0, \"Input shape not entirely divisible by number of patches\"\n        assert chw[2] % n_patches == 0, \"Input shape not entirely divisible by number of patches\"\n        self.patch_size = (chw[1] / n_patches, chw[2] / n_patches)\n\n        # 1) Linear mapper\n        self.input_d = int(chw[0] * self.patch_size[0] * self.patch_size[1])\n        self.linear_mapper = nn.Linear(self.input_d, self.hidden_d)\n        \n        # 2) Learnable classification token\n        self.class_token = nn.Parameter(torch.rand(1, self.hidden_d))\n        \n        # 3) Positional embedding\n        self.pos_embed = nn.Parameter(get_positional_embeddings(self.n_patches ** 2 + 1, self.hidden_d).clone())\n        self.pos_embed.requires_grad = False\n        \n        # 4) Transformer encoder blocks\n        self.blocks = nn.ModuleList([MyViTBlock(hidden_d, n_heads) for _ in range(n_blocks)])\n        \n        # 5) Classification MLPk\n        self.mlp = nn.Sequential(\n            nn.Linear(self.hidden_d, out_d),\n            nn.Softmax(dim=-1)\n        )\n\n    def forward(self, images):\n        # Dividing images into patches\n        n, c, h, w = images.shape\n        patches = patchify(images, self.n_patches).to(self.pos_embed.device)\n        \n        # Running linear layer tokenization\n        # Map the vector corresponding to each patch to the hidden size dimension\n        tokens = self.linear_mapper(patches)\n        \n        # Adding classification token to the tokens\n        tokens = torch.stack([torch.vstack((self.class_token, tokens[i])) for i in range(len(tokens))])\n        \n        # Adding positional embedding\n        pos_embed = self.pos_embed.repeat(n, 1, 1)\n        out = tokens + pos_embed\n        \n        # Transformer Blocks\n        for block in self.blocks:\n            out = block(out)\n            \n        # Getting the classification token only\n        out = out[:, 0]\n        \n        return self.mlp(out) # Map to output dimension, output category distribution\n    \n\n\ndef get_positional_embeddings(sequence_length, d):\n    result = torch.ones(sequence_length, d)\n    for i in range(sequence_length):\n        for j in range(d):\n            result[i][j] = np.sin(i / (10000 ** (j / d))) if j % 2 == 0 else np.cos(i / (10000 ** ((j - 1) / d)))\n    return result","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:56:39.665761Z","iopub.execute_input":"2022-09-15T08:56:39.666716Z","iopub.status.idle":"2022-09-15T08:56:39.695644Z","shell.execute_reply.started":"2022-09-15T08:56:39.666679Z","shell.execute_reply":"2022-09-15T08:56:39.694588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main():\n    # Get the train and valid data\n    fold=2\n\n    train = df[df[\"fold\"] != fold].reset_index(drop=True)\n    valid = df[df[\"fold\"] == fold].reset_index(drop=True)\n    \n    train_dataset = RSNADataset(csv=train, mode=\"train\", \n                                transform=train_transforms)\n    valid_dataset = RSNADataset(csv=valid, mode=\"train\", \n                                transform=valid_transforms)\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    # Defining model and training options\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(\"Using device: \", device, f\"({torch.cuda.get_device_name(device)})\" if torch.cuda.is_available() else \"\")\n    model = MyViT((3, 28, 28), n_patches=7, n_blocks=2, hidden_d=8, n_heads=2, out_d=8).to(device)\n    N_EPOCHS = 5\n    LR = 0.005\n\n    # Training loop\n    optimizer = Adam(model.parameters(), lr=LR)\n    criterion = CrossEntropyLoss()\n    for epoch in tqdm(range(N_EPOCHS), desc=\"Training\"):\n        train_loss = 0.0\n        for data in tqdm(trainloader, desc=f\"Epoch {epoch + 1} in training\", leave=False):\n            image, targets = data_to_device(data)\n            img=torch.mean(image, -1)\n            img=img.permute(0,1,2,3)\n            y_hat = model(img)\n            loss = criterion(y_hat, targets)\n\n            train_loss += loss.detach().cpu().item() / len(train_loader)\n\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n        print(f\"Epoch {epoch + 1}/{N_EPOCHS} loss: {train_loss:.2f}\")\n\n    # Test loop\n    with torch.no_grad():\n        correct, total = 0, 0\n        test_loss = 0.0\n        for data in tqdm(validloader, desc=\"Testing\"):\n            image, targets = data_to_device(data)\n            img=torch.mean(image, -1)\n            img=img.permute(0,1,2,3)\n            y_hat = model(img)\n            loss = criterion(y_hat, targets)\n            test_loss += loss.detach().cpu().item() / len(test_loader)\n\n            correct += torch.sum(torch.argmax(y_hat, dim=1) == y).detach().cpu().item()\n            total += len(x)\n        print(f\"Test loss: {test_loss:.2f}\")\n        print(f\"Test accuracy: {correct / total * 100:.2f}%\")\n\n\nif __name__ == '__main__':\n    main()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T08:56:39.699011Z","iopub.execute_input":"2022-09-15T08:56:39.699775Z","iopub.status.idle":"2022-09-15T09:01:17.794164Z","shell.execute_reply.started":"2022-09-15T08:56:39.699672Z","shell.execute_reply":"2022-09-15T09:01:17.792514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}