{"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":"# ViT for RSNA 2023\n\n[ViT original github](https://github.com/lucidrains/vit-pytorch), more specifically the [ViT 3D](https://github.com/lucidrains/vit-pytorch/blob/main/vit_pytorch/vit_3d.py). \n\nApply on 3D data.\n\nAttempt to train on the 2 GPU T4. Not working for now, please comment if you see why !\n\n---\n\n# ViT 3D implementation\n\n","metadata":{}},{"cell_type":"code","source":"!pip install einops","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:44:57.354373Z","iopub.execute_input":"2023-10-12T12:44:57.354765Z","iopub.status.idle":"2023-10-12T12:45:09.809479Z","shell.execute_reply.started":"2023-10-12T12:44:57.354741Z","shell.execute_reply":"2023-10-12T12:45:09.808386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install lightning","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:45:09.811649Z","iopub.execute_input":"2023-10-12T12:45:09.812579Z","iopub.status.idle":"2023-10-12T12:45:26.851515Z","shell.execute_reply.started":"2023-10-12T12:45:09.812543Z","shell.execute_reply":"2023-10-12T12:45:26.85045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pathlib import Path\n\nimport torch\nfrom torch import nn\nfrom torch.nn.utils.rnn import pad_sequence\nfrom torch.nn.functional import cross_entropy\nfrom torch.utils.data import Dataset, random_split\nimport torch.optim as optim\n\nfrom einops import rearrange, repeat\nfrom einops.layers.torch import Rearrange\n\nimport polars as pl\nfrom sklearn.model_selection import train_test_split\nimport numpy as np\n\nfrom datetime import datetime\nfrom torch.utils.tensorboard import SummaryWriter\n\nimport math\n\n# for parallelization\nimport lightning.pytorch as li","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:45:26.853306Z","iopub.execute_input":"2023-10-12T12:45:26.853672Z","iopub.status.idle":"2023-10-12T12:45:52.224039Z","shell.execute_reply.started":"2023-10-12T12:45:26.853637Z","shell.execute_reply":"2023-10-12T12:45:52.223144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Incompatible with Lightning\n#device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:45:52.225861Z","iopub.execute_input":"2023-10-12T12:45:52.22613Z","iopub.status.idle":"2023-10-12T12:45:52.231366Z","shell.execute_reply.started":"2023-10-12T12:45:52.226108Z","shell.execute_reply":"2023-10-12T12:45:52.230575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_PATH = Path(\"/kaggle/input/rsna-2023-abdominal-trauma-detection\")\nPT_PATH = Path(\"/kaggle/input/rsna-2023-5mm-slices-pt\")\nOUTPUT_PATH = Path(\"/kaggle/working\")\n\nLABEL_COLS = [\"bowel_healthy\", \"bowel_injury\", \"extravasation_healthy\", \"extravasation_injury\", \"kidney_healthy\", \"kidney_low\", \"kidney_high\", \"liver_healthy\", \"liver_low\", \"liver_high\", \"spleen_healthy\", \"spleen_low\", \"spleen_high\", \"any_injury\"]\nFRAME_PATCH_SIZE = 4\nBATCH_SIZE = 2","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:45:52.234707Z","iopub.execute_input":"2023-10-12T12:45:52.235343Z","iopub.status.idle":"2023-10-12T12:45:52.250438Z","shell.execute_reply.started":"2023-10-12T12:45:52.235323Z","shell.execute_reply":"2023-10-12T12:45:52.249568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"[The very first brick of our model](https://pytorch.org/tutorials/beginner/introyt/modelsyt_tutorial.html) is a python class inheriting from torch.nn.Module \n\n>One important behavior of torch.nn.Module is registering parameters. If a particular Module subclass has learning weights, these weights are expressed as instances of torch.nn.Parameter. The Parameter class is a subclass of torch.Tensor, with the special behavior that when they are assigned as attributes of a Module, they are added to the list of that modules parameters. These parameters may be accessed through the parameters() method on the Module class.","metadata":{}},{"cell_type":"code","source":"# helpers\n\ndef pair(t):\n    return t if isinstance(t, tuple) else (t, t)\n\n# classes\n\nclass FeedForward(nn.Module):\n    def __init__(self, dim, hidden_dim, dropout = 0.):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.LayerNorm(dim),\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.norm = nn.LayerNorm(dim)\n        self.attend = nn.Softmax(dim = -1)\n        self.dropout = nn.Dropout(dropout)\n\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        x = self.norm(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        attn = self.dropout(attn)\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        for _ in range(depth):\n            self.layers.append(nn.ModuleList([\n                Attention(dim, heads = heads, dim_head = dim_head, dropout = dropout),\n                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\n\nclass ViT(nn.Module):\n    def __init__(self, *, image_size, image_patch_size, frames, frame_patch_size, num_classes, dim, depth, heads, mlp_dim, pool = 'cls', channels = 3, dim_head = 64, dropout = 0., emb_dropout = 0.):\n        super().__init__()\n        image_height, image_width = pair(image_size)\n        patch_height, patch_width = pair(image_patch_size)\n\n        assert image_height % patch_height == 0 and image_width % patch_width == 0, 'Image dimensions must be divisible by the patch size.'\n        assert frames % frame_patch_size == 0, 'Frames must be divisible by frame patch size'\n\n        num_patches = (image_height // patch_height) * (image_width // patch_width) * (frames // frame_patch_size)\n        patch_dim = channels * patch_height * patch_width * frame_patch_size\n\n        assert pool in {'cls', 'mean'}, 'pool type must be either cls (cls token) or mean (mean pooling)'\n\n        self.to_patch_embedding = nn.Sequential(\n            Rearrange('b c (f pf) (h p1) (w p2) -> b (f h w) (p1 p2 pf c)', p1 = patch_height, p2 = patch_width, pf = frame_patch_size),\n            nn.LayerNorm(patch_dim),\n            nn.Linear(patch_dim, dim),\n            nn.LayerNorm(dim),\n        )\n\n        self.pos_embedding = nn.Parameter(torch.randn(1, num_patches + 1, dim))\n        self.cls_token = nn.Parameter(torch.randn(1, 1, dim))\n        self.dropout = nn.Dropout(emb_dropout)\n\n        self.transformer = Transformer(dim, depth, heads, dim_head, mlp_dim, dropout)\n\n        self.pool = pool\n        self.to_latent = nn.Identity()\n\n        self.mlp_head = nn.Sequential(\n            nn.LayerNorm(dim),\n            nn.Linear(dim, num_classes)\n        )\n\n    def forward(self, video):\n        x = self.to_patch_embedding(video)\n        b, n, _ = x.shape\n\n        cls_tokens = repeat(self.cls_token, '1 1 d -> b 1 d', b = b)\n        x = torch.cat((cls_tokens, x), dim=1)\n        x += self.pos_embedding[:, :(n + 1)]\n        x = self.dropout(x)\n\n        x = self.transformer(x)\n\n        x = x.mean(dim = 1) if self.pool == 'mean' else x[:, 0]\n\n        x = self.to_latent(x)\n        return self.mlp_head(x)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:46:16.828139Z","iopub.execute_input":"2023-10-12T12:46:16.828481Z","iopub.status.idle":"2023-10-12T12:46:16.847636Z","shell.execute_reply.started":"2023-10-12T12:46:16.828455Z","shell.execute_reply":"2023-10-12T12:46:16.846666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ViT(\n    image_size = 128,          # image size\n    frames = 256,               # number of frames\n    image_patch_size = 16,     # image patch size\n    frame_patch_size = FRAME_PATCH_SIZE,      # frame patch size\n    num_classes = len(LABEL_COLS[:-1]),\n    dim = 512,\n    depth = 6,\n    heads = 4,\n    mlp_dim = 1024,\n    channels = 1,\n    dropout = 0.1,\n    emb_dropout = 0.1\n)\n\n# Incompatible with Lightning\n# model.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:46:17.36799Z","iopub.execute_input":"2023-10-12T12:46:17.368349Z","iopub.status.idle":"2023-10-12T12:46:17.556429Z","shell.execute_reply.started":"2023-10-12T12:46:17.368321Z","shell.execute_reply":"2023-10-12T12:46:17.555424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#video = torch.randn(2, 1, 256, 128, 128) # (batch, channels, frames, height, width)\n# Incompatible with Lightning\n#video.to(device)\n#video.shape","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:46:18.22239Z","iopub.execute_input":"2023-10-12T12:46:18.223439Z","iopub.status.idle":"2023-10-12T12:46:18.227862Z","shell.execute_reply.started":"2023-10-12T12:46:18.223396Z","shell.execute_reply":"2023-10-12T12:46:18.226914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#video.dtype","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:46:18.449375Z","iopub.execute_input":"2023-10-12T12:46:18.45046Z","iopub.status.idle":"2023-10-12T12:46:18.454928Z","shell.execute_reply.started":"2023-10-12T12:46:18.45042Z","shell.execute_reply":"2023-10-12T12:46:18.453966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#preds = model(video)\n#preds","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:46:18.89566Z","iopub.execute_input":"2023-10-12T12:46:18.896663Z","iopub.status.idle":"2023-10-12T12:46:18.900768Z","shell.execute_reply.started":"2023-10-12T12:46:18.896597Z","shell.execute_reply":"2023-10-12T12:46:18.899811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data preparation","metadata":{}},{"cell_type":"code","source":"train_labels = pl.read_csv(BASE_PATH.joinpath(\"train.csv\"))\ntrain_patient_series = pl.read_csv(BASE_PATH.joinpath(\"train_series_meta.csv\"))\n\ncounts_unique_labels = train_labels[LABEL_COLS].group_by(LABEL_COLS).count().with_row_count(offset=1)\ncounts_unique_labels = counts_unique_labels.with_columns(pl.when(pl.col(\"count\") == 1).then(0).otherwise(pl.col(\"row_nr\")).alias(\"group\"))\ncounts_unique_labels = counts_unique_labels.with_columns(pl.when(pl.col(\"count\") == 2292).then((len(counts_unique_labels)-1)/pl.col(\"count\")).otherwise(1/pl.col(\"count\")).alias(\"weight\"))\ntrain_labels_group = train_labels.join(counts_unique_labels.select(LABEL_COLS + [\"group\", \"count\", \"weight\"]), on=LABEL_COLS, how=\"left\")","metadata":{"execution":{"iopub.status.busy":"2023-10-12T14:06:48.823723Z","iopub.execute_input":"2023-10-12T14:06:48.824199Z","iopub.status.idle":"2023-10-12T14:06:48.849723Z","shell.execute_reply.started":"2023-10-12T14:06:48.824152Z","shell.execute_reply":"2023-10-12T14:06:48.848739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels_group","metadata":{"execution":{"iopub.status.busy":"2023-10-12T14:06:49.689469Z","iopub.execute_input":"2023-10-12T14:06:49.689895Z","iopub.status.idle":"2023-10-12T14:06:49.70346Z","shell.execute_reply.started":"2023-10-12T14:06:49.689844Z","shell.execute_reply":"2023-10-12T14:06:49.702443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_idx, valid_idx= train_test_split(\nnp.arange(len(train_labels)),\ntest_size=0.1,\nshuffle=True,\nstratify=train_labels_group[\"group\"])","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:46:21.110674Z","iopub.execute_input":"2023-10-12T12:46:21.111009Z","iopub.status.idle":"2023-10-12T12:46:21.134876Z","shell.execute_reply.started":"2023-10-12T12:46:21.110985Z","shell.execute_reply":"2023-10-12T12:46:21.133906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RNSA2023Dataset(Dataset):\n    def __init__(self, labels_data, idx):\n        self.labels = labels_data[idx]\n        self.patients_series = pl.read_csv(BASE_PATH.joinpath(\"train_series_meta.csv\")).join(self.labels.select(\"patient_id\"), on=\"patient_id\", how=\"inner\")\n        \n    def __len__(self):\n        return len(self.patients_series)\n    \n    def __getitem__(self, idx):\n        patient, serie, _, _ = self.patients_series[idx]\n        data = torch.load(PT_PATH.joinpath(str(patient[0])).joinpath(f\"{str(serie[0])}.pt\"))#, map_location=torch.device(device)) # Incompatible with Lightning\n        # Do not consider any_injury from LABEL_COLS as a label\n        label = torch.from_numpy(self.labels.filter(pl.col(\"patient_id\")==patient[0]).select(pl.col(LABEL_COLS[:-1])).to_numpy())#.to(device) # Incompatible with Lightning\n        return data, label","metadata":{"execution":{"iopub.status.busy":"2023-10-12T14:49:52.007495Z","iopub.execute_input":"2023-10-12T14:49:52.007919Z","iopub.status.idle":"2023-10-12T14:49:52.019993Z","shell.execute_reply.started":"2023-10-12T14:49:52.007884Z","shell.execute_reply":"2023-10-12T14:49:52.018897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = RNSA2023Dataset(train_labels_group, train_idx)\nvalid_dataset = RNSA2023Dataset(train_labels_group, valid_idx)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T14:50:09.367486Z","iopub.execute_input":"2023-10-12T14:50:09.367828Z","iopub.status.idle":"2023-10-12T14:50:09.387367Z","shell.execute_reply.started":"2023-10-12T14:50:09.3678Z","shell.execute_reply":"2023-10-12T14:50:09.386339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#train_dataset[0] # works only if using device in the cells above, commented to use Lightning","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:46:29.092463Z","iopub.execute_input":"2023-10-12T12:46:29.092936Z","iopub.status.idle":"2023-10-12T12:46:29.097472Z","shell.execute_reply.started":"2023-10-12T12:46:29.092897Z","shell.execute_reply":"2023-10-12T12:46:29.096609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_patient_series), len(train_dataset), len(valid_dataset), len(train_dataset) + len(valid_dataset) == len(train_patient_series)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:46:29.324161Z","iopub.execute_input":"2023-10-12T12:46:29.324577Z","iopub.status.idle":"2023-10-12T12:46:29.334174Z","shell.execute_reply.started":"2023-10-12T12:46:29.324527Z","shell.execute_reply":"2023-10-12T12:46:29.333289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Adapted from https://discuss.pytorch.org/t/dataloader-for-various-length-of-data/6418/18\ndef collate_fn_pad(items):\n    data_seq = []\n    label_seq = []\n    max_frame = 0 \n    for data, label in items:\n        if data.shape[0] > max_frame:\n            max_frame = data.shape[0]\n        data_seq.append(data)\n        label_seq.append(label)\n    big_one = torch.ones((FRAME_PATCH_SIZE*(math.floor(max_frame/FRAME_PATCH_SIZE)+1), data.shape[1], data.shape[2]))\n    data_seq.append(big_one)\n    data_seq_batched = pad_sequence(data_seq, batch_first=True).unsqueeze(1)[:-1] # unsqueeze to add channel, do not keep the big one\n    label_seq_batched = torch.cat(label_seq)\n    assert data_seq_batched.shape[0] == len(label_seq_batched)\n    return data_seq_batched.to(torch.float32), label_seq_batched.to(torch.float32)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T12:46:30.187765Z","iopub.execute_input":"2023-10-12T12:46:30.188163Z","iopub.status.idle":"2023-10-12T12:46:30.199957Z","shell.execute_reply.started":"2023-10-12T12:46:30.188132Z","shell.execute_reply":"2023-10-12T12:46:30.198924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create data loaders for our datasets; shuffle for training, not for validation\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, collate_fn = collate_fn_pad)\nvalid_loader = torch.utils.data.DataLoader(valid_dataset, batch_size=BATCH_SIZE, shuffle=False, collate_fn = collate_fn_pad)","metadata":{"execution":{"iopub.status.busy":"2023-10-03T16:48:17.700402Z","iopub.execute_input":"2023-10-03T16:48:17.700715Z","iopub.status.idle":"2023-10-03T16:48:17.718373Z","shell.execute_reply.started":"2023-10-03T16:48:17.700687Z","shell.execute_reply":"2023-10-03T16:48:17.717593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"# Define the LightningModule\nclass LitViT(li.LightningModule):\n    def __init__(self, vit):\n        super().__init__()\n        self.vit = vit\n        \n    def forward(self, x):\n        return self.vit(x)\n    \n    def training_step(self, batch, batch_idx):\n        # training_step defines the train loop.\n        # it is independent of forward\n        inputs, labels = batch\n        outputs = self.vit(inputs)\n        loss = cross_entropy(outputs, labels)\n        self.log(\"train_loss\", loss)\n        return loss\n    \n    def configure_optimizers(self):\n        optimizer = optim.Adam(self.parameters(), lr=1e-3)\n        return optimizer","metadata":{"execution":{"iopub.status.busy":"2023-10-03T16:48:17.749021Z","iopub.execute_input":"2023-10-03T16:48:17.750982Z","iopub.status.idle":"2023-10-03T16:48:17.761645Z","shell.execute_reply.started":"2023-10-03T16:48:17.750952Z","shell.execute_reply":"2023-10-03T16:48:17.760878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Init the LitVit\nlit_vit = LitViT(model)","metadata":{"execution":{"iopub.status.busy":"2023-10-03T16:48:17.762868Z","iopub.execute_input":"2023-10-03T16:48:17.76338Z","iopub.status.idle":"2023-10-03T16:48:17.777019Z","shell.execute_reply.started":"2023-10-03T16:48:17.763351Z","shell.execute_reply":"2023-10-03T16:48:17.776228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Init trainer\ntrainer = li.Trainer(devices=\"auto\", \n                     accelerator=\"gpu\", \n                     strategy='ddp_notebook', \n                     max_epochs=1, \n                     log_every_n_steps=1)","metadata":{"execution":{"iopub.status.busy":"2023-10-03T16:48:17.778294Z","iopub.execute_input":"2023-10-03T16:48:17.778886Z","iopub.status.idle":"2023-10-03T16:48:18.095156Z","shell.execute_reply.started":"2023-10-03T16:48:17.778851Z","shell.execute_reply":"2023-10-03T16:48:18.093833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train the model (hint: here are some helpful Trainer arguments for rapid idea iteration)\ntrainer.fit(model=lit_vit, train_dataloaders=train_loader)","metadata":{"execution":{"iopub.status.busy":"2023-10-03T16:48:18.096531Z","iopub.execute_input":"2023-10-03T16:48:18.09749Z","iopub.status.idle":"2023-10-03T16:56:23.838344Z","shell.execute_reply.started":"2023-10-03T16:48:18.097454Z","shell.execute_reply":"2023-10-03T16:56:23.836923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---\n\n---","metadata":{}},{"cell_type":"code","source":"'''\ndef train_one_epoch(epoch_index, tb_writer):\n    running_loss = 0.\n    last_loss = 0.\n\n    # Here, we use enumerate(training_loader) instead of\n    # iter(training_loader) so that we can track the batch\n    # index and do some intra-epoch reporting\n    for i, data in enumerate(train_loader):\n        # Every data instance is an input + label pair\n        inputs, labels = data\n\n        # Zero your gradients for every batch!\n        optimizer.zero_grad()\n\n        # Make predictions for this batch\n        outputs = model(inputs)\n\n        # Compute the loss and its gradients\n        loss = loss_fn(outputs, labels)\n        loss.backward()\n\n        # Adjust learning weights\n        optimizer.step()\n\n        # Gather data and report\n        running_loss += loss.item()\n        if i % 1000 == 999:\n            last_loss = running_loss / 1000 # loss per batch\n            print('  batch {} loss: {}'.format(i + 1, last_loss))\n            tb_x = epoch_index * len(training_loader) + i + 1\n            tb_writer.add_scalar('Loss/train', last_loss, tb_x)\n            running_loss = 0.\n\n    return last_loss\n'''","metadata":{"execution":{"iopub.status.busy":"2023-10-03T16:56:23.84078Z","iopub.execute_input":"2023-10-03T16:56:23.841878Z","iopub.status.idle":"2023-10-03T16:56:23.850453Z","shell.execute_reply.started":"2023-10-03T16:56:23.841832Z","shell.execute_reply":"2023-10-03T16:56:23.849584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\n# Initializing in a separate cell so we can easily add more epochs to the same run\ntimestamp = datetime.now().strftime('%Y%m%d_%H%M%S')\nwriter = SummaryWriter(OUTPUT_PATH.joinpath(f\"train_{timestamp}\"))\nepoch_number = 0\n\nEPOCHS = 5\n\nbest_vloss = 1_000_000.\n\nfor epoch in range(EPOCHS):\n    print('EPOCH {}:'.format(epoch_number + 1))\n\n    # Make sure gradient tracking is on, and do a pass over the data\n    model.train(True)\n    avg_loss = train_one_epoch(epoch_number, writer)\n\n\n    running_vloss = 0.0\n    # Set the model to evaluation mode, disabling dropout and using population\n    # statistics for batch normalization.\n    model.eval()\n\n    # Disable gradient computation and reduce memory consumption.\n    with torch.no_grad():\n        for i, vdata in enumerate(valid_loader):\n            vinputs, vlabels = vdata\n            voutputs = model(vinputs)\n            vloss = loss_fn(voutputs, vlabels)\n            running_vloss += vloss\n\n    avg_vloss = running_vloss / (i + 1)\n    print('LOSS train {} valid {}'.format(avg_loss, avg_vloss))\n\n    # Log the running loss averaged per batch\n    # for both training and validation\n    writer.add_scalars('Training vs. Validation Loss',\n                    { 'Training' : avg_loss, 'Validation' : avg_vloss },\n                    epoch_number + 1)\n    writer.flush()\n\n    # Track best performance, and save the model's state\n    if avg_vloss < best_vloss:\n        best_vloss = avg_vloss\n        model_path = 'model_{}_{}'.format(timestamp, epoch_number)\n        torch.save(model.state_dict(), model_path)\n\n    epoch_number += 1\n'''","metadata":{"execution":{"iopub.status.busy":"2023-10-03T16:56:23.85208Z","iopub.execute_input":"2023-10-03T16:56:23.852735Z","iopub.status.idle":"2023-10-03T16:56:23.866449Z","shell.execute_reply.started":"2023-10-03T16:56:23.852705Z","shell.execute_reply":"2023-10-03T16:56:23.865576Z"},"trusted":true},"execution_count":null,"outputs":[]}]}