{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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":8066583,"sourceType":"datasetVersion","datasetId":4732842}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Hello fellow Kagglers,\n\nThis notebook demonstrates the training strategy using EfficientVit and is a work in progress.\n\nInference notebook and further updates will follow soon!\n\n[dataset](https://www.kaggle.com/code/markwijkhuizen/birdclef-2024-eda-preprocessed-dataset)","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport imageio.v3 as imageio\nimport albumentations as A\nimport matplotlib.pyplot as plt\n\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm.notebook import tqdm\nfrom torchvision import transforms\n\nimport torch\nimport torchmetrics\nimport timm\nimport pickle\nimport psutil\nimport time\nimport os","metadata":{"execution":{"iopub.status.busy":"2024-04-08T18:44:10.793247Z","iopub.execute_input":"2024-04-08T18:44:10.794091Z","iopub.status.idle":"2024-04-08T18:44:23.345594Z","shell.execute_reply.started":"2024-04-08T18:44:10.794058Z","shell.execute_reply":"2024-04-08T18:44:23.344779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"class Config:\n    # Dataset\n    HEIGHT = 128\n    WIDTH = 320\n    ROOT_FOLDER = '/kaggle/input/birdclef-2024-dataset'\n    # Training\n    BATCH_SIZE = 16\n    N_EPOCHS = 20\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    \nCONFIG = Config()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T18:44:23.347399Z","iopub.execute_input":"2024-04-08T18:44:23.347745Z","iopub.status.idle":"2024-04-08T18:44:23.353916Z","shell.execute_reply.started":"2024-04-08T18:44:23.347701Z","shell.execute_reply":"2024-04-08T18:44:23.352976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Sample Submission","metadata":{}},{"cell_type":"code","source":"sample_submission = pd.read_csv('/kaggle/input/birdclef-2024/sample_submission.csv')\n\n# Set labels\nCONFIG.LABELS = sample_submission.columns[1:]\nCONFIG.N_CLASSES = len(CONFIG.LABELS)\nprint(f'# classes: {CONFIG.N_CLASSES}')\n\ndisplay(sample_submission.head())","metadata":{"execution":{"iopub.status.busy":"2024-04-08T18:44:23.355046Z","iopub.execute_input":"2024-04-08T18:44:23.355315Z","iopub.status.idle":"2024-04-08T18:44:23.408169Z","shell.execute_reply.started":"2024-04-08T18:44:23.355283Z","shell.execute_reply":"2024-04-08T18:44:23.407310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"code","source":"# Load Spectrogram PNG Bytes\nwith open(f'{CONFIG.ROOT_FOLDER}/X.pkl', 'rb') as file:\n    X = pickle.load(file)\n    \n# Load Labels\nwith open(f'{CONFIG.ROOT_FOLDER}/y.pkl', 'rb') as file:\n    y = 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\nprint(f'N_SAMPLES: {CONFIG.N_SAMPLES:,}')\nprint(f'N_STEPS_PER_EPOCH: {CONFIG.N_STEPS_PER_EPOCH:,}')\nprint(f'N_STEPS: {CONFIG.N_STEPS:,}')","metadata":{"execution":{"iopub.status.busy":"2024-04-08T18:44:23.410388Z","iopub.execute_input":"2024-04-08T18:44:23.410645Z","iopub.status.idle":"2024-04-08T18:44:54.158155Z","shell.execute_reply.started":"2024-04-08T18:44:23.410623Z","shell.execute_reply":"2024-04-08T18:44:54.157162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"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-04-08T18:44:54.159375Z","iopub.execute_input":"2024-04-08T18:44:54.159664Z","iopub.status.idle":"2024-04-08T18:44:54.165199Z","shell.execute_reply.started":"2024-04-08T18:44:54.159637Z","shell.execute_reply":"2024-04-08T18:44:54.164140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"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-04-08T18:44:54.166298Z","iopub.execute_input":"2024-04-08T18:44:54.166601Z","iopub.status.idle":"2024-04-08T18:44:54.178668Z","shell.execute_reply.started":"2024-04-08T18:44:54.166577Z","shell.execute_reply":"2024-04-08T18:44:54.177858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train\ntrain_dataset = MyDataset(X, y, TRAIN_TRANSFORMS)\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    )\ntrain_dataloader_iter = iter(train_dataloader)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T18:44:54.179774Z","iopub.execute_input":"2024-04-08T18:44:54.180030Z","iopub.status.idle":"2024-04-08T18:44:54.420280Z","shell.execute_reply.started":"2024-04-08T18:44:54.180008Z","shell.execute_reply":"2024-04-08T18:44:54.418106Z"},"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-04-08T18:44:54.423353Z","iopub.execute_input":"2024-04-08T18:44:54.424660Z","iopub.status.idle":"2024-04-08T18:44:54.702646Z","shell.execute_reply.started":"2024-04-08T18:44:54.424619Z","shell.execute_reply":"2024-04-08T18:44:54.701810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Plot Batch","metadata":{}},{"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-04-08T18:44:54.706440Z","iopub.execute_input":"2024-04-08T18:44:54.708284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# 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('efficientvit_b1')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Count model parameters\ndef count_parameters(model):\n    return sum([p.numel() for p in model.parameters()])","metadata":{"execution":{"iopub.status.idle":"2024-04-08T18:44:58.706842Z"},"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    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":{"execution":{"iopub.status.busy":"2024-04-08T18:44:58.707944Z","iopub.execute_input":"2024-04-08T18:44:58.708199Z","iopub.status.idle":"2024-04-08T18:44:58.718881Z"},"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":{"execution":{"iopub.status.busy":"2024-04-08T18:44:58.720027Z","iopub.execute_input":"2024-04-08T18:44:58.720278Z","iopub.status.idle":"2024-04-08T18:45:00.410152Z","shell.execute_reply.started":"2024-04-08T18:44:58.720256Z","shell.execute_reply":"2024-04-08T18:45:00.409151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Learning Rate Scheduler","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2024-04-08T18:45:00.411328Z","iopub.execute_input":"2024-04-08T18:45:00.411613Z","iopub.status.idle":"2024-04-08T18:45:00.416864Z","shell.execute_reply.started":"2024-04-08T18:45:00.411589Z","shell.execute_reply":"2024-04-08T18:45:00.415923Z"},"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":{"execution":{"iopub.status.busy":"2024-04-08T18:45:00.418371Z","iopub.execute_input":"2024-04-08T18:45:00.418715Z","iopub.status.idle":"2024-04-08T18:45:00.953257Z","shell.execute_reply.started":"2024-04-08T18:45:00.418681Z","shell.execute_reply":"2024-04-08T18:45:00.952312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"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-04-08T18:45:00.954642Z","iopub.execute_input":"2024-04-08T18:45:00.955019Z","iopub.status.idle":"2024-04-08T18:45:00.961041Z","shell.execute_reply.started":"2024-04-08T18:45:00.954975Z","shell.execute_reply":"2024-04-08T18:45:00.960059Z"},"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-04-08T18:45:00.962193Z","iopub.execute_input":"2024-04-08T18:45:00.962479Z","iopub.status.idle":"2024-04-08T18:45:00.975535Z","shell.execute_reply.started":"2024-04-08T18:45:00.962455Z","shell.execute_reply":"2024-04-08T18:45:00.974552Z"},"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\n# Save entire model object\ntorch.save(model, 'model.pth')","metadata":{"execution":{"iopub.status.busy":"2024-04-08T18:45:00.976491Z","iopub.execute_input":"2024-04-08T18:45:00.976940Z"},"trusted":true},"execution_count":null,"outputs":[]}]}