{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":227680449,"sourceType":"kernelVersion"},{"sourceId":229103448,"sourceType":"kernelVersion"}],"dockerImageVersionId":30918,"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 torch\nimport torchaudio\nfrom joblib import Parallel, delayed\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nimport random\nimport timm\n\ntorch.manual_seed(42);","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T01:44:03.988131Z","iopub.execute_input":"2025-03-23T01:44:03.988416Z","iopub.status.idle":"2025-03-23T01:44:03.994136Z","shell.execute_reply.started":"2025-03-23T01:44:03.988395Z","shell.execute_reply":"2025-03-23T01:44:03.993432Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nDEVICE","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T01:44:03.995201Z","iopub.execute_input":"2025-03-23T01:44:03.995399Z","iopub.status.idle":"2025-03-23T01:44:04.014137Z","shell.execute_reply.started":"2025-03-23T01:44:03.995382Z","shell.execute_reply":"2025-03-23T01:44:04.013298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_parquet('/kaggle/input/bc25-eda/train_metadata_joined.parquet')\n\nSPEC_CROP = 248\n\nIDX_TO_LABEL = sorted(df.primary_label.unique())\nLABEL_TO_IDX = dict((v, k) for k, v in enumerate(IDX_TO_LABEL))\nNUM_LABELS = len(IDX_TO_LABEL)\nLABEL_TO_ONEHOT = torch.eye(NUM_LABELS, device=DEVICE)\nLABELS = torch.tensor([LABEL_TO_IDX[l] for l in df.primary_label.iloc], device=DEVICE)\nMIN_SPEC_VALUES = torch.tensor(df.min_spec_value, dtype=torch.float32, device=DEVICE)\nMAX_SPEC_VALUES = torch.tensor(df.max_spec_value, dtype=torch.float32, device=DEVICE)\n\nNUM_SAMPLES = len(df)\nsplit = int(0.8 * NUM_SAMPLES)\nall_indices = torch.randperm(NUM_SAMPLES)\nTRAIN_INDICES = all_indices[:split]\nVAL_INDICES = all_indices[split:]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T01:44:04.015545Z","iopub.execute_input":"2025-03-23T01:44:04.015775Z","iopub.status.idle":"2025-03-23T01:44:04.230888Z","shell.execute_reply.started":"2025-03-23T01:44:04.015755Z","shell.execute_reply":"2025-03-23T01:44:04.230191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_file(filename):\n    filename = '/kaggle/input/bc25-train-specs-signed-db-80-n-16/train_audio_specs/' + filename.split('.')[0] + '.pt'\n    return torch.load(filename, weights_only=True, map_location=DEVICE)\n\nALL_SPECS = Parallel(n_jobs=-1)(\n    delayed(load_file)(fname) for fname in tqdm(df.filename, desc=\"Loading files\")\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T01:44:04.231925Z","iopub.execute_input":"2025-03-23T01:44:04.232150Z","iopub.status.idle":"2025-03-23T01:45:23.072847Z","shell.execute_reply.started":"2025-03-23T01:44:04.232130Z","shell.execute_reply":"2025-03-23T01:45:23.072149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Quantizer:\n    def __init__(self, num_bits):\n        self.range = 2**num_bits\n        self.max = 2**(num_bits - 1) - 1\n        self.min = -2**(num_bits - 1)\n        if num_bits <= 8:\n            self.dtype = torch.int8\n        elif num_bits <= 16:\n            self.dtype = torch.int16\n        elif num_bits <= 32:\n            self.dtype = torch.int32\n\n    def quantize(self, tensor):\n        min_val = tensor.min()\n        max_val = tensor.max()\n        if min_val == max_val: # Edge case: all values are the same\n            return torch.full_like(tensor, 0, dtype=self.dtype), min_val, max_val\n        scale = self.range / (max_val - min_val)\n        quantized_tensor = torch.round((tensor - min_val) * scale + self.min).clamp(self.min, self.max).to(self.dtype)\n        return quantized_tensor, min_val, max_val\n\n    def dequantize(self, quantized_tensor, min_val, max_val):\n        if min_val == max_val:\n            return torch.full_like(quantized_tensor, min_val, dtype=torch.float32)\n        scale = (max_val - min_val) / self.range\n        return (quantized_tensor.to(torch.float32) - self.min) * scale + min_val\n\nq = Quantizer(16)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T01:45:23.073729Z","iopub.execute_input":"2025-03-23T01:45:23.073976Z","iopub.status.idle":"2025-03-23T01:45:23.080822Z","shell.execute_reply.started":"2025-03-23T01:45:23.073955Z","shell.execute_reply":"2025-03-23T01:45:23.080036Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataloader","metadata":{}},{"cell_type":"code","source":"class SpecDataset(torch.utils.data.Dataset):\n    def __init__(self, validation):\n        self.validation = validation\n    def __getitem__(self, i):\n        spec = ALL_SPECS[i]\n        max_range = spec.shape[1] - SPEC_CROP\n        if self.validation:\n            spec = spec[:, :SPEC_CROP] # only first 5 sec\n        else:\n            start = random.randint(0, max(0, max_range))\n            spec = spec[:, start:start+SPEC_CROP] # random crop\n        if max_range < 0:\n            spec = torch.nn.functional.pad(spec, (0, -max_range)) # fill with zeros\n        spec = q.dequantize(spec, MIN_SPEC_VALUES[i], MAX_SPEC_VALUES[i])\n        return spec[None], LABEL_TO_ONEHOT[LABELS[i]]\n    def __len__(self):\n        return len(df)\n\n\ndef make_dataloaders():\n    train_dataset = torch.utils.data.Subset(SpecDataset(validation=False), TRAIN_INDICES)\n    val_dataset = torch.utils.data.Subset(SpecDataset(validation=True), VAL_INDICES)\n    train_dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=16, shuffle=True)\n    val_dataloader = torch.utils.data.DataLoader(val_dataset, batch_size=16, shuffle=False)\n    return train_dataloader, val_dataloader\n\n\nbatch = next(iter(make_dataloaders()[1]))\nplt.imshow(batch[0][0][0].to('cpu'))\nplt.colorbar();","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T01:45:23.082344Z","iopub.execute_input":"2025-03-23T01:45:23.082589Z","iopub.status.idle":"2025-03-23T01:45:23.379764Z","shell.execute_reply.started":"2025-03-23T01:45:23.082569Z","shell.execute_reply":"2025-03-23T01:45:23.378788Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"def make_model():\n    return timm.create_model(\n        'tf_efficientnet_b0',\n        in_chans=1,\n        num_classes=NUM_LABELS,\n        pretrained=False,\n    )\n\nmake_model()(batch[0].to('cpu')).shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T01:45:23.380848Z","iopub.execute_input":"2025-03-23T01:45:23.381083Z","iopub.status.idle":"2025-03-23T01:45:24.002285Z","shell.execute_reply.started":"2025-03-23T01:45:23.381064Z","shell.execute_reply":"2025-03-23T01:45:24.001312Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"def get_val_metrics(logits, labels, lossfn):\n    loss = lossfn(logits, labels).item()\n    acc = (labels == (logits > 0.0)).float().mean().item()\n    return dict(loss=loss, acc=acc)\n\n\ndef train(model, optimizer, num_epochs, train_losses=[], val_metrics=[]):\n    train_dataloader, val_dataloader = make_dataloaders()\n    lossfn = torch.nn.BCEWithLogitsLoss()\n    for epoch in tqdm(range(num_epochs)):\n        model.train()\n        for specs, labels in train_dataloader:\n            optimizer.zero_grad()\n            logits = model(specs)\n            loss = lossfn(logits, labels)\n            loss.backward()\n            optimizer.step()\n            train_losses.append(loss.detach().item())\n        model.eval()\n        with torch.inference_mode():\n            epoch_labels = []\n            epoch_logits = []\n            for specs, labels in val_dataloader:\n                logits = model(specs)\n                epoch_labels.append(labels)\n                epoch_logits.append(logits)\n            epoch_labels = torch.concat(epoch_labels)\n            epoch_logits = torch.concat(epoch_logits)\n            val_metrics.append(get_val_metrics(epoch_logits, epoch_labels, lossfn))\n        torch.save(model.state_dict(), f'model_state_dict_epoch_{epoch}.pt')\n    return train_losses, val_metrics\n\ndef plot_losses(train_losses, val_metrics):\n    num_train_batches = len(make_dataloaders()[0])\n    plt.plot(train_losses)\n    plt.plot([v['loss'] for v in val_metrics for _ in range(num_train_batches)])\n\n\n\nmodel = make_model().to(DEVICE)\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\ntrain_losses, val_metrics = train(model, optimizer, 20)\nplot_losses(train_losses, val_metrics)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T01:45:24.003226Z","iopub.execute_input":"2025-03-23T01:45:24.003491Z","execution_failed":"2025-03-23T01:46:44.856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-23T01:46:44.856Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation","metadata":{}}]}