{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":12162756,"sourceType":"datasetVersion","datasetId":7660181}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport torch.nn as nn\nimport librosa\nimport os\nimport torchaudio\nimport torch\nfrom IPython.display import Audio, display\nimport timm\nfrom torch.utils.data import Dataset,DataLoader,random_split","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:34:15.444904Z","iopub.execute_input":"2025-07-01T09:34:15.445073Z","iopub.status.idle":"2025-07-01T09:34:15.449612Z","shell.execute_reply.started":"2025-07-01T09:34:15.445057Z","shell.execute_reply":"2025-07-01T09:34:15.448945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:34:53.679452Z","iopub.execute_input":"2025-07-01T09:34:53.680207Z","iopub.status.idle":"2025-07-01T09:34:53.792687Z","shell.execute_reply.started":"2025-07-01T09:34:53.680169Z","shell.execute_reply":"2025-07-01T09:34:53.791850Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!nvidia-smi","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:34:53.917864Z","iopub.execute_input":"2025-07-01T09:34:53.918128Z","iopub.status.idle":"2025-07-01T09:34:54.161314Z","shell.execute_reply.started":"2025-07-01T09:34:53.918108Z","shell.execute_reply":"2025-07-01T09:34:54.160325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformers import AutoFeatureExtractor, AutoModelForAudioClassification,AutoConfig\n\n\ntotol_target_class = 206\nconfig = AutoConfig.from_pretrained(\"Simon-Kotchou/ssast-small-patch-audioset-16-16\")\nfeature_extractor = AutoFeatureExtractor.from_pretrained(\"Simon-Kotchou/ssast-small-patch-audioset-16-16\")\nmodel = AutoModelForAudioClassification.from_pretrained(\"Simon-Kotchou/ssast-small-patch-audioset-16-16\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:34:54.163287Z","iopub.execute_input":"2025-07-01T09:34:54.163645Z","iopub.status.idle":"2025-07-01T09:35:15.535312Z","shell.execute_reply.started":"2025-07-01T09:34:54.163597Z","shell.execute_reply":"2025-07-01T09:35:15.534469Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(model.classifier)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:15.536778Z","iopub.execute_input":"2025-07-01T09:35:15.537483Z","iopub.status.idle":"2025-07-01T09:35:15.541712Z","shell.execute_reply.started":"2025-07-01T09:35:15.537453Z","shell.execute_reply":"2025-07-01T09:35:15.540865Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.classifier = nn.Sequential(\n    nn.LayerNorm(384,eps=1e-12),\n    nn.Linear(384,206)\n    \n)\nconfig.num_labels = 206","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:15.542497Z","iopub.execute_input":"2025-07-01T09:35:15.542817Z","iopub.status.idle":"2025-07-01T09:35:15.568412Z","shell.execute_reply.started":"2025-07-01T09:35:15.542787Z","shell.execute_reply":"2025-07-01T09:35:15.567819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(model.classifier)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:15.570269Z","iopub.execute_input":"2025-07-01T09:35:15.570496Z","iopub.status.idle":"2025-07-01T09:35:15.583144Z","shell.execute_reply.started":"2025-07-01T09:35:15.570478Z","shell.execute_reply":"2025-07-01T09:35:15.582423Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = torch.nn.DataParallel(model)  # use all available GPUs\nmodel = model.to(device)\n\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\nloss_fn = torch.nn.CrossEntropyLoss()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:15.583906Z","iopub.execute_input":"2025-07-01T09:35:15.584101Z","iopub.status.idle":"2025-07-01T09:35:15.827830Z","shell.execute_reply.started":"2025-07-01T09:35:15.584086Z","shell.execute_reply":"2025-07-01T09:35:15.826949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_function(waveform,max_duration=5):\n    \n    inputs = feature_extractor(\n        waveform,\n        sampling_rate=feature_extractor.sampling_rate,\n        max_length=int(feature_extractor.sampling_rate * max_duration),\n        truncation=True,\n        return_attention_mask=True,\n        return_tensors=\"pt\",\n    )\n    return inputs\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:15.829171Z","iopub.execute_input":"2025-07-01T09:35:15.829436Z","iopub.status.idle":"2025-07-01T09:35:15.833541Z","shell.execute_reply.started":"2025-07-01T09:35:15.829408Z","shell.execute_reply":"2025-07-01T09:35:15.832661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Dataset\nimport os\nimport librosa\n\nclass CustomDataset(Dataset):\n    def __init__(self, audio_folder, preprocess_function, class_names=None, class_to_idx=None):\n        super(CustomDataset, self).__init__()\n        self.audio_folder = audio_folder\n        self.clipped_audio = []\n        self.audio_labels = []\n        self.preprocess_function = preprocess_function\n\n        if class_names is None:\n            self.class_names = sorted(os.listdir(self.audio_folder))\n        else:\n            self.class_names = class_names\n\n        if class_to_idx is None:\n            self.class_to_idx = {class_name: idx for idx, class_name in enumerate(self.class_names)}\n        else:\n            self.class_to_idx = class_to_idx\n\n        for class_name in self.class_names:\n            class_path = os.path.join(self.audio_folder, class_name)\n            for audio_file in os.listdir(class_path):\n                self.clipped_audio.append(os.path.join(class_path, audio_file))\n                self.audio_labels.append(self.class_to_idx[class_name])\n\n    def __len__(self):\n        return len(self.audio_labels)\n\n    def __getitem__(self, idx):\n        audio_clip = self.clipped_audio[idx]\n        label = self.audio_labels[idx]\n        waveform, sample_rate = librosa.load(audio_clip, sr=None)\n        model_input = self.preprocess_function(waveform) # returns dictionary\n        model_input = {k: torch.tensor(v).squeeze(0) if isinstance(v, (list, np.ndarray)) else v.squeeze(0) for k, v in model_input.items()}\n        \n        return model_input,label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:15.834550Z","iopub.execute_input":"2025-07-01T09:35:15.834986Z","iopub.status.idle":"2025-07-01T09:35:15.872599Z","shell.execute_reply.started":"2025-07-01T09:35:15.834956Z","shell.execute_reply":"2025-07-01T09:35:15.872015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"audio_folder = \"/kaggle/input/preprocessed-data/kaggle/working/output\"\ndataset = CustomDataset(audio_folder,preprocess_function)\nclass_name =   dataset.class_names\nclass_to_idx = dataset.class_to_idx\n\nidx_to_class = {v:k for k,v in class_to_idx.items()}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:15.873402Z","iopub.execute_input":"2025-07-01T09:35:15.873673Z","iopub.status.idle":"2025-07-01T09:35:18.802079Z","shell.execute_reply.started":"2025-07-01T09:35:15.873657Z","shell.execute_reply":"2025-07-01T09:35:18.801253Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"total_size = len(dataset)\nval_size = int(0.1*total_size)\ntrain_size = total_size- val_size\n\ntrain_dataset,valid_dataset = random_split(dataset,[train_size,val_size])\n\ntrain_dataloader = DataLoader(train_dataset,batch_size=32,shuffle=True)\nvalid_dataloader = DataLoader(valid_dataset,batch_size=32,shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:18.802994Z","iopub.execute_input":"2025-07-01T09:35:18.803272Z","iopub.status.idle":"2025-07-01T09:35:18.811387Z","shell.execute_reply.started":"2025-07-01T09:35:18.803249Z","shell.execute_reply":"2025-07-01T09:35:18.810640Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i,(audio,label) in enumerate(train_dataloader):\n    print(f\"audio shape : {audio['input_values'].shape}\")\n    break\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:18.814480Z","iopub.execute_input":"2025-07-01T09:35:18.814844Z","iopub.status.idle":"2025-07-01T09:35:33.202343Z","shell.execute_reply.started":"2025-07-01T09:35:18.814817Z","shell.execute_reply":"2025-07-01T09:35:33.201648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Epochs = 25\ntrain_loss_list = []\nvalid_loss_list = []\nbest_loss = 1_000_000\nbest_model = None\n\nfor epoch in range(Epochs):\n    model.train()\n    training_loss = 0.0\n    validation_loss = 0.0\n\n    # Training loop\n    for i, (audio, label) in enumerate(train_dataloader):\n        audio = {k: v.to(device) for k, v in audio.items()}  # move each tensor to device\n        label = label.to(device)\n    \n        optimizer.zero_grad()\n        output = model(**audio)\n        loss = loss_fn(output.logits, label)  \n        loss.backward()\n        optimizer.step()\n        training_loss += loss.item()\n\n    # Validation loop\n    model.eval()\n    with torch.no_grad():\n        for j, (vaudio, vlabel) in enumerate(valid_dataloader):\n            vaudio = {k: v.to(device) for k, v in vaudio.items()}\n            vlabel = vlabel.to(device)\n    \n            voutput = model(**vaudio)\n            vloss = loss_fn(voutput.logits, vlabel)\n            validation_loss += vloss.item()\n\n    \n    print(f\"Epoch [{epoch+1}/{Epochs}] - Train Loss: {training_loss:.4f}, Val Loss: {validation_loss:.4f}\")\n    train_loss_list.append(training_loss/(i+1))\n    valid_loss_list.append(validation_loss/(j+1))\n\n    if validation_loss < best_loss:\n        best_loss = validation_loss\n        best_model = model\n        torch.save(model.state_dict(), \"best_model.pth\")\n\n    \n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:35:33.203863Z","iopub.execute_input":"2025-07-01T09:35:33.204346Z","iopub.status.idle":"2025-07-01T14:42:11.437627Z","shell.execute_reply.started":"2025-07-01T09:35:33.204327Z","shell.execute_reply":"2025-07-01T14:42:11.436760Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Assuming best_model is wrapped in DataParallel\n# And you want to save Hugging Face-compatible model, config, and feature extractor\n\nsave_dir = \"ssast-206-final\"\n\n# 1. Save the actual model (not the DataParallel wrapper)\nbest_model.module.save_pretrained(save_dir)\n\n# 2. Save the feature extractor (if you're using AutoFeatureExtractor or similar)\nfeature_extractor.save_pretrained(save_dir)\n\n# 3. Save the model config (optional if not already included)\nconfig.save_pretrained(save_dir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T14:53:34.687817Z","iopub.execute_input":"2025-07-01T14:53:34.688073Z","iopub.status.idle":"2025-07-01T14:53:34.833764Z","shell.execute_reply.started":"2025-07-01T14:53:34.688055Z","shell.execute_reply":"2025-07-01T14:53:34.833132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport shutil\nshutil.make_archive(save_dir, 'zip', save_dir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T14:54:40.993178Z","iopub.execute_input":"2025-07-01T14:54:40.993761Z","iopub.status.idle":"2025-07-01T14:54:45.469153Z","shell.execute_reply.started":"2025-07-01T14:54:40.993734Z","shell.execute_reply":"2025-07-01T14:54:45.468340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplt.figure(figsize=(12,8))\nplt.plot(train_loss_list, label=\"Train Loss\")\nplt.plot(valid_loss_list, label=\"Validation Loss\")\nplt.title(\"Loss vs Epochs\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Loss\")\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T14:53:47.576289Z","iopub.execute_input":"2025-07-01T14:53:47.576911Z","iopub.status.idle":"2025-07-01T14:53:47.883258Z","shell.execute_reply.started":"2025-07-01T14:53:47.576880Z","shell.execute_reply":"2025-07-01T14:53:47.882485Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}