{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":182571486,"sourceType":"kernelVersion"}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport imageio.v3 as imageio\nimport albumentations as A\n\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm.notebook import tqdm\nfrom torchvision import transforms\nfrom torchvision.models import wide_resnet50_2, Wide_ResNet50_2_Weights\n\nimport torch\nimport torchmetrics\nimport timm\nimport psutil\nimport time\n\nimport librosa\nimport cv2\nimport pickle\nimport lzma\nimport os","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:36:54.898097Z","iopub.execute_input":"2024-06-10T17:36:54.898981Z","iopub.status.idle":"2024-06-10T17:36:54.905350Z","shell.execute_reply.started":"2024-06-10T17:36:54.898947Z","shell.execute_reply":"2024-06-10T17:36:54.904200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config():\n\n    # Competition Root Folder\n    ROOT_FOLDER= '/kaggle/input/birdclef-2024'\n    DATA_FOLDER = '/kaggle/input/birdclef-2024-dataset-preparation'\n\n    # Dataset\n    HEIGHT = 128\n    WIDTH = 320\n\n    # Training\n    BATCH_SIZE = 16\n    N_EPOCHS = 25\n    # Model\n    BACKBONE = 'efficientvit_b1.r288_in1k'\n    # Learning Rate Scheduler\n    LR_MAX = 3e-4\n    WEIGHT_DECAY = 0.00\n    # Others\n    SEED = 42\n    # IS_INTERACTIVE = os.environ['KAGGLE_KERNEL_RUN_TYPE'] == 'Interactive'\n    IS_INTERACTIVE = True\n    \nCONFIG = Config()","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:36:54.906693Z","iopub.execute_input":"2024-06-10T17:36:54.906941Z","iopub.status.idle":"2024-06-10T17:36:54.916579Z","shell.execute_reply.started":"2024-06-10T17:36:54.906919Z","shell.execute_reply":"2024-06-10T17:36:54.915669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission = pd.read_csv(f'{CONFIG.ROOT_FOLDER}/sample_submission.csv')\n\n# Set labels\nCONFIG.LABELS = sample_submission.columns[1:]\nCONFIG.N_LABELS = len(CONFIG.LABELS)\nCONFIG.N_CLASSES = len(CONFIG.LABELS)","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:36:54.926529Z","iopub.execute_input":"2024-06-10T17:36:54.926801Z","iopub.status.idle":"2024-06-10T17:36:54.939904Z","shell.execute_reply.started":"2024-06-10T17:36:54.926777Z","shell.execute_reply":"2024-06-10T17:36:54.939160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load Spectrogram PNG Bytes\nwith open(f'{CONFIG.DATA_FOLDER}/X_train.pkl', 'rb') as file:\n    X = pickle.load(file)\n    \n# Load Labels\nwith open(f'{CONFIG.DATA_FOLDER}/y_train.pkl', 'rb') as file:\n    y = pickle.load(file)\n    \n\n# Load Spectrogram PNG Bytes\nwith open(f'{CONFIG.DATA_FOLDER}/X_val.pkl', 'rb') as file:\n    X_val = pickle.load(file)\n    \n# Load Labels\nwith open(f'{CONFIG.DATA_FOLDER}/y_val.pkl', 'rb') as file:\n    y_val = pickle.load(file)\n    \nCONFIG.N_SAMPLES = len(y)\nCONFIG.N_STEPS_PER_EPOCH = CONFIG.N_SAMPLES // CONFIG.BATCH_SIZE\nCONFIG.N_STEPS = CONFIG.N_STEPS_PER_EPOCH * CONFIG.N_EPOCHS\n\nCONFIG.N_VAL_SAMPLES = len(y_val)\nCONFIG.N_VAL_STEPS_PER_EPOCH = CONFIG.N_VAL_SAMPLES // CONFIG.BATCH_SIZE\nCONFIG.N_VAL_STEPS = CONFIG.N_VAL_STEPS_PER_EPOCH * CONFIG.N_EPOCHS\n\n\nprint(\n    f'CONFIG.N_SAMPLES: {CONFIG.N_SAMPLES}\\n'\n    f'CONFIG.N_STEPS_PER_EPOCH: {CONFIG.N_STEPS_PER_EPOCH}\\n'\n    f'CONFIG.N_STEPS: {CONFIG.N_STEPS}\\n'\n    f'CONFIG.N_VAL_SAMPLES: {CONFIG.N_VAL_SAMPLES}\\n'\n    f'CONFIG.N_VAL_STEPS_PER_EPOCH: {CONFIG.N_VAL_STEPS_PER_EPOCH}\\n'\n    f'CONFIG.N_VAL_STEPS: {CONFIG.N_VAL_STEPS}\\n'\n)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:36:54.955898Z","iopub.execute_input":"2024-06-10T17:36:54.956251Z","iopub.status.idle":"2024-06-10T17:37:22.314226Z","shell.execute_reply.started":"2024-06-10T17:36:54.956218Z","shell.execute_reply":"2024-06-10T17:37:22.313213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training Augmentations\nTRAIN_TRANSFORMS = A.Compose([\n        A.RandomBrightnessContrast(brightness_limit=0.10, contrast_limit=0.10, p=0.50),\n        A.ImageCompression(quality_lower=75, quality_upper=100, p=0.5),\n    ])","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:22.316353Z","iopub.execute_input":"2024-06-10T17:37:22.316672Z","iopub.status.idle":"2024-06-10T17:37:22.322521Z","shell.execute_reply.started":"2024-06-10T17:37:22.316643Z","shell.execute_reply":"2024-06-10T17:37:22.321447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self, X, y, transforms=None):\n        self.X = X\n        self.y = y\n        self.keys = tuple(X.keys())\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.X)\n\n    def __getitem__(self, index):\n        # Acquire Key\n        key = self.keys[index]\n        # Read Decode PNG Spectrogram\n        spec = imageio.imread(self.X[key])\n        std_global = spec.std()\n        # Random Offset\n        _, W = spec.shape\n        if W < CONFIG.WIDTH: # Pad\n            spec = np.pad(spec, ((0,0), (0,CONFIG.WIDTH-W)))\n        elif W > CONFIG.WIDTH: # Crop\n            offset = np.random.randint(0, W-CONFIG.WIDTH)\n            spec = spec[:,offset:offset+CONFIG.WIDTH]\n        \n        label = self.y[key]\n        \n        return spec, label","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:22.324108Z","iopub.execute_input":"2024-06-10T17:37:22.325042Z","iopub.status.idle":"2024-06-10T17:37:22.335838Z","shell.execute_reply.started":"2024-06-10T17:37:22.325001Z","shell.execute_reply":"2024-06-10T17:37:22.334814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Data loaders\ntrain_dataset = MyDataset(X, y, TRAIN_TRANSFORMS)\nval_dataset = MyDataset(X_val, y_val, TRAIN_TRANSFORMS)\n\nval_dataloader = DataLoader(\n        val_dataset,\n        batch_size=CONFIG.BATCH_SIZE,\n        shuffle=True,\n        drop_last=True,\n        num_workers=psutil.cpu_count(),\n        # num_workers=0,\n    )\n\ntrain_dataloader = DataLoader(\n        train_dataset,\n        batch_size=CONFIG.BATCH_SIZE,\n        shuffle=True,\n        drop_last=True,\n        num_workers=psutil.cpu_count(),\n        # num_workers=0,\n    )\n\nval_dataloader_iter = iter(val_dataloader)\ntrain_dataloader_iter = iter(train_dataloader)","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:22.337288Z","iopub.execute_input":"2024-06-10T17:37:22.337645Z","iopub.status.idle":"2024-06-10T17:37:22.942357Z","shell.execute_reply.started":"2024-06-10T17:37:22.337613Z","shell.execute_reply":"2024-06-10T17:37:22.939885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Example batch\nX_batch, y_batch = next(train_dataloader_iter)\n# X_batch\nprint(f'X_batch shape: {X_batch.shape}, dtype: {X_batch.dtype}')\nprint(f'X_batch min: {X_batch.min():.3f}, max: {X_batch.max():.3f}')\nprint(f'X_batch µ: {X_batch.float().mean():.3f}, σ: {X_batch.float().std():.3f}')\n# Label\nprint(f'y_batch shape: {y_batch.shape}, dtype: {y_batch.dtype}')\nprint(f'y_batch min: {y_batch.min()}, max: {y_batch.max()}')","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:22.947568Z","iopub.execute_input":"2024-06-10T17:37:22.948115Z","iopub.status.idle":"2024-06-10T17:37:23.158344Z","shell.execute_reply.started":"2024-06-10T17:37:22.948046Z","shell.execute_reply":"2024-06-10T17:37:23.157258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot a training batch\ndef plot_batch(nrows=8, ncols=2):\n    fig, axes = plt.subplots(nrows, ncols, figsize=(ncols*6, nrows*4))\n    for r in range(nrows):\n        for c in range(ncols):\n            idx = (r * ncols) + c\n            # Denormalize Image\n            axes[r,c].imshow(X_batch[idx])\n            axes[r,c].set_title(f'shape: {X_batch[idx].numpy().shape}, label: {y_batch[idx]}')\n    plt.show()\n    \nplot_batch()","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:23.159757Z","iopub.execute_input":"2024-06-10T17:37:23.160804Z","iopub.status.idle":"2024-06-10T17:37:27.268096Z","shell.execute_reply.started":"2024-06-10T17:37:23.160767Z","shell.execute_reply":"2024-06-10T17:37:27.266706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#TODO: REMOVE\n# Function to search for timm models\ndef search_timm_model(query):\n    search_result = [n for n in timm.list_models(pretrained=True) if query in n]\n    for i, name in enumerate(search_result):\n        print(f'{i:02d} | {name}')\n        \nsearch_timm_model('efficientnet_b0')","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:27.269505Z","iopub.execute_input":"2024-06-10T17:37:27.269847Z","iopub.status.idle":"2024-06-10T17:37:27.287880Z","shell.execute_reply.started":"2024-06-10T17:37:27.269807Z","shell.execute_reply":"2024-06-10T17:37:27.286996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#TODO: Check and remove\n# Count model parameters\ndef count_parameters(model):\n    return sum([p.numel() for p in model.parameters()])","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:27.289135Z","iopub.execute_input":"2024-06-10T17:37:27.289486Z","iopub.status.idle":"2024-06-10T17:37:27.298880Z","shell.execute_reply.started":"2024-06-10T17:37:27.289463Z","shell.execute_reply":"2024-06-10T17:37:27.297937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # ImageNet Normalize Input\n        self.normalize = transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n        \n        # Backbone\n        self.backbone = timm.create_model(\n                CONFIG.BACKBONE,\n                pretrained=True,\n                num_classes=CONFIG.N_CLASSES,\n            )\n        \n        \n    def forward(self, inputs):\n        # Go From HxW → 3xHxW\n        inputs = inputs.unsqueeze(1).expand(-1, 3, -1, -1)\n        # Normalize [0-255] → [0-1]\n        inputs = inputs.float() / 255\n        # Normalize\n        inputs = self.normalize(inputs)\n        \n        return self.backbone(inputs)","metadata":{"ExecuteTime":{"end_time":"2024-06-08T18:20:18.394597Z","start_time":"2024-06-08T18:20:18.384642Z"},"execution":{"iopub.status.busy":"2024-06-10T17:37:27.300048Z","iopub.execute_input":"2024-06-10T17:37:27.300366Z","iopub.status.idle":"2024-06-10T17:37:27.312694Z","shell.execute_reply.started":"2024-06-10T17:37:27.300341Z","shell.execute_reply":"2024-06-10T17:37:27.310531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create new Model\nmodel = Model().cuda()\n\n# Number of parameters\nprint(f'# Model Parameters: {count_parameters(model):,}')\n\n# Forward pass\nwith torch.no_grad():\n    # Put inputs on GPU\n    outputs = model(X_batch.cuda())\n    print(f'outputs shape: {outputs.shape}, min: {outputs.min():.3f}, max: {outputs.max():.3f}')\n    print(f'µ: {outputs.mean():.3f}, σ: {outputs.std():.3f}')","metadata":{"ExecuteTime":{"end_time":"2024-06-08T18:22:21.221749Z","start_time":"2024-06-08T18:22:20.851371Z"},"execution":{"iopub.status.busy":"2024-06-10T17:37:27.314594Z","iopub.execute_input":"2024-06-10T17:37:27.315060Z","iopub.status.idle":"2024-06-10T17:37:31.057520Z","shell.execute_reply.started":"2024-06-10T17:37:27.315022Z","shell.execute_reply":"2024-06-10T17:37:31.056559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get the learning rate scheduler\ndef get_lr_scheduler(optimizer):\n    return torch.optim.lr_scheduler.OneCycleLR(\n        optimizer=optimizer,\n        max_lr=CONFIG.LR_MAX,\n        total_steps=CONFIG.N_STEPS,\n        pct_start=0.10,\n        anneal_strategy='cos',\n        div_factor=1e3,\n        final_div_factor=1e4,\n    )","metadata":{"ExecuteTime":{"end_time":"2024-06-08T18:23:26.938532Z","start_time":"2024-06-08T18:23:26.933909Z"},"execution":{"iopub.status.busy":"2024-06-10T17:37:31.058800Z","iopub.execute_input":"2024-06-10T17:37:31.059111Z","iopub.status.idle":"2024-06-10T17:37:31.064560Z","shell.execute_reply.started":"2024-06-10T17:37:31.059086Z","shell.execute_reply":"2024-06-10T17:37:31.063560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot Learning Rate Scheduler\ndef plot_lr_scheduler():\n    lr_scheduler = get_lr_scheduler(torch.optim.Adam(model.parameters()))\n    lrs  = []\n    for step in range(CONFIG.N_STEPS):\n        lrs.append(lr_scheduler.get_last_lr())\n        lr_scheduler.step()\n    # Plot Learning Rate\n    plt.figure(figsize=(12,5))\n    plt.title('Learning Rate Schedule')\n    plt.xticks(np.arange(0, CONFIG.N_STEPS+1, CONFIG.N_STEPS_PER_EPOCH), range(CONFIG.N_EPOCHS+1))\n    plt.xlim(0, CONFIG.N_STEPS)\n    plt.ylim(0, CONFIG.LR_MAX*1.1)\n    plt.xlabel('Epoch')\n    plt.ylabel('Learning Rate')\n    plt.plot(lrs)\n    plt.grid()\n    plt.show()\n    # Reset Learning Rate Scheduler\n    lr_scheduler._step_count = 0\n    lr_scheduler.last_epoch = 0\n\nplot_lr_scheduler()","metadata":{"ExecuteTime":{"end_time":"2024-06-08T18:23:28.852698Z","start_time":"2024-06-08T18:23:28.793070Z"},"execution":{"iopub.status.busy":"2024-06-10T17:37:31.065734Z","iopub.execute_input":"2024-06-10T17:37:31.066039Z","iopub.status.idle":"2024-06-10T17:37:31.729335Z","shell.execute_reply.started":"2024-06-10T17:37:31.066014Z","shell.execute_reply":"2024-06-10T17:37:31.728364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Average meter to keep track of metrics/loss during training\nclass AverageMeter(object):\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val):\n        self.sum += val.sum()\n        self.count += val.numel()\n        # Average is simply the sum divided by the count\n        self.avg = self.sum / self.count","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:31.730564Z","iopub.execute_input":"2024-06-10T17:37:31.730863Z","iopub.status.idle":"2024-06-10T17:37:31.736973Z","shell.execute_reply.started":"2024-06-10T17:37:31.730831Z","shell.execute_reply":"2024-06-10T17:37:31.736111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Loss\nloss_fn = nn.CrossEntropyLoss()\n# Optimizer\noptimizer = torch.optim.AdamW(\n    params=model.parameters(),\n    lr=CONFIG.LR_MAX,\n    weight_decay=CONFIG.WEIGHT_DECAY,\n)\n# Learning Rate Scheduler\nLR_SCHEDULER = get_lr_scheduler(optimizer)\n# Metrics\nLOSS = AverageMeter()\nACC = torchmetrics.Accuracy(task='multiclass', num_classes=CONFIG.N_CLASSES).cuda()\nROC_AUC = torchmetrics.AUROC(task='multiclass', num_classes=CONFIG.N_CLASSES).cuda()","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:31.741468Z","iopub.execute_input":"2024-06-10T17:37:31.742055Z","iopub.status.idle":"2024-06-10T17:37:31.751586Z","shell.execute_reply.started":"2024-06-10T17:37:31.742026Z","shell.execute_reply":"2024-06-10T17:37:31.750804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Initialize Metrics Lists\ntrain_losses = []\nval_losses = []\ntrain_accuracies = []\nval_accuracies = []","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:31.752778Z","iopub.execute_input":"2024-06-10T17:37:31.753077Z","iopub.status.idle":"2024-06-10T17:37:31.758788Z","shell.execute_reply.started":"2024-06-10T17:37:31.753052Z","shell.execute_reply":"2024-06-10T17:37:31.757877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(CONFIG.N_EPOCHS):\n    # Reset Metrics\n    LOSS.reset()\n    ACC.reset()\n    ROC_AUC.reset()\n    # Put model in training mode\n    model.train()\n    # Iterate Over Training Dataloader\n    for step, (X_batch, y_true) in enumerate(train_dataloader):\n        # Put batch on GPU\n        X_batch = X_batch.cuda()\n        y_true = y_true.cuda()\n        # Step Time\n        t_start = time.perf_counter()\n        # Forward Pass\n        y_pred = model(X_batch)\n        # Loss\n        loss = loss_fn(y_pred, y_true)\n        # Update Loss Metrics\n        LOSS.update(loss)\n        # Compute Gradients\n        loss.backward()\n        # Backward Pass\n        optimizer.step()\n        # Zero Out Gradients\n        optimizer.zero_grad()\n        # Update Metrics\n        ACC.update(y_pred.softmax(dim=1), y_true)\n        ROC_AUC.update(y_pred.softmax(dim=1), y_true)\n        # Logs\n        if not CONFIG.IS_INTERACTIVE and (step + 1) == CONFIG.N_STEPS_PER_EPOCH:\n            print(\n                f'EPOCH {epoch+1:02d} {step+1:04d}/{CONFIG.N_STEPS_PER_EPOCH} | ' +\n                f'loss: {LOSS.avg:.4f}, ACC: {ACC.compute():.3f}, ROC_AUC: {ROC_AUC.compute():.3f}, ' +\n                f'step: {(time.perf_counter()-t_start):.3f}s, lr: {LR_SCHEDULER.get_last_lr()[0]:.2e}',\n            )\n        elif CONFIG.IS_INTERACTIVE:\n            print(\n                f'EPOCH {epoch+1:02d} {step+1:04d}/{CONFIG.N_STEPS_PER_EPOCH} | ' +\n                f'loss: {LOSS.avg:.4f}, ACC: {ACC.compute():.3f}, ROC_AUC: {ROC_AUC.compute():.3f}, ' +\n                f'step: {(time.perf_counter()-t_start):.3f}s, lr: {LR_SCHEDULER.get_last_lr()[0]:.2e}'\n                , end='\\n' if (step + 1) == CONFIG.N_STEPS_PER_EPOCH else ' ' * 10 + '\\r', flush=True,\n            )\n        # Learning Rate Scheduler Step\n        LR_SCHEDULER.step()\n    # Save Metrics\n    train_losses.append(LOSS.avg.item())\n    train_accuracies.append(ACC.compute().item())\n\n    # Validation phase\n    # Put model in evaluation mode\n    model.eval()\n    # Iterate Over Validation Dataloader\n    for step, (X_batch, y_true) in enumerate(val_dataloader):\n        # Put batch on GPU\n        X_batch = X_batch.cuda()\n        y_true = y_true.cuda()\n        # Forward Pass\n        with torch.no_grad():\n            y_pred = model(X_batch)\n        # Loss\n        loss = loss_fn(y_pred, y_true)\n        # Update Loss Metrics\n        LOSS.update(loss)\n        # Update Metrics\n        ACC.update(y_pred.softmax(dim=1), y_true)\n        ROC_AUC.update(y_pred.softmax(dim=1), y_true)\n        # Logs\n        if not CONFIG.IS_INTERACTIVE and (step + 1) == CONFIG.N_VAL_STEPS_PER_EPOCH:\n            print(\n                f'VAL {step+1:04d}/{CONFIG.N_VAL_STEPS_PER_EPOCH} | ' +\n                f'loss: {LOSS.avg:.4f}, ACC: {ACC.compute():.3f}, ROC_AUC: {ROC_AUC.compute():.3f}',\n            )\n        elif CONFIG.IS_INTERACTIVE:\n            print(\n                f'VAL {step+1:04d}/{CONFIG.N_VAL_STEPS_PER_EPOCH} | ' +\n                f'loss: {LOSS.avg:.4f}, ACC: {ACC.compute():.3f}, ROC_AUC: {ROC_AUC.compute():.3f}'\n                , end='\\n' if (step + 1) == CONFIG.N_VAL_STEPS_PER_EPOCH else ' ' * 10 + '\\r', flush=True,\n            )\n    # Save Metrics\n    val_losses.append(LOSS.avg.item())\n    val_accuracies.append(ACC.compute().item())\n\n# Save entire model object\ntorch.save(model, 'model_split_data.pth')","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:37:31.760608Z","iopub.execute_input":"2024-06-10T17:37:31.760997Z","iopub.status.idle":"2024-06-10T17:41:51.128615Z","shell.execute_reply.started":"2024-06-10T17:37:31.760963Z","shell.execute_reply":"2024-06-10T17:41:51.127194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# Plot training and validation loss\nplt.figure(figsize=(12, 6))\nplt.subplot(1, 2, 1)\nplt.plot(train_losses, label='Training Loss')\nplt.plot(val_losses, label='Validation Loss')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.legend()\n\n# Plot training and validation accuracy\nplt.subplot(1, 2, 2)\nplt.plot(train_accuracies, label='Training Accuracy')\nplt.plot(val_accuracies, label='Validation Accuracy')\nplt.xlabel('Epochs')\nplt.ylabel('Accuracy')\nplt.legend()\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-10T17:41:51.129881Z","iopub.status.idle":"2024-06-10T17:41:51.130284Z","shell.execute_reply.started":"2024-06-10T17:41:51.130087Z","shell.execute_reply":"2024-06-10T17:41:51.130104Z"},"trusted":true},"execution_count":null,"outputs":[]}]}