{"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":11053663,"sourceType":"datasetVersion","datasetId":6886569}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport logging\nimport random\nimport gc\nimport time\nimport math\nimport warnings\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nimport librosa\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.auto import tqdm\n\nfrom transformers import ASTForAudioClassification, Wav2Vec2FeatureExtractor, get_scheduler\nfrom transformers import ASTFeatureExtractor\nimport math\n\nwarnings.filterwarnings(\"ignore\")\nlogging.basicConfig(level=logging.ERROR)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:57:49.582134Z","iopub.execute_input":"2025-04-04T13:57:49.582523Z","iopub.status.idle":"2025-04-04T13:58:09.721313Z","shell.execute_reply.started":"2025-04-04T13:57:49.582493Z","shell.execute_reply":"2025-04-04T13:58:09.720332Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    debug = False\n    print_freq = 100\n    num_workers = 0\n    \n    OUTPUT_DIR = '/kaggle/working/'\n    train_datadir = '/kaggle/input/birdclef-2025/train_audio'\n    train_csv = '/kaggle/input/birdclef-2025/train.csv'\n    taxonomy_csv = '/kaggle/input/birdclef-2025/taxonomy.csv'\n    \n    model_name = 'MIT/ast-finetuned-audioset-10-10-0.4593'\n    pretrained = True\n    \n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    epochs = 3\n    batch_size = 16\n    criterion = 'BCEWithLogitsLoss'\n    \n    optimizer = 'AdamW'\n    lr = 3e-5\n    weight_decay = 1e-5\n    train_perc = 0.7","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:09.722644Z","iopub.execute_input":"2025-04-04T13:58:09.723320Z","iopub.status.idle":"2025-04-04T13:58:09.780738Z","shell.execute_reply.started":"2025-04-04T13:58:09.723288Z","shell.execute_reply":"2025-04-04T13:58:09.779776Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cfg = CFG()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:09.782817Z","iopub.execute_input":"2025-04-04T13:58:09.783086Z","iopub.status.idle":"2025-04-04T13:58:09.799808Z","shell.execute_reply.started":"2025-04-04T13:58:09.783064Z","shell.execute_reply":"2025-04-04T13:58:09.799006Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(cfg.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:09.801140Z","iopub.execute_input":"2025-04-04T13:58:09.801422Z","iopub.status.idle":"2025-04-04T13:58:09.821199Z","shell.execute_reply.started":"2025-04-04T13:58:09.801395Z","shell.execute_reply":"2025-04-04T13:58:09.820490Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"feature_extractor = ASTFeatureExtractor.from_pretrained(cfg.model_name)\n#feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(cfg.model_name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:09.821999Z","iopub.execute_input":"2025-04-04T13:58:09.822196Z","iopub.status.idle":"2025-04-04T13:58:10.146523Z","shell.execute_reply.started":"2025-04-04T13:58:09.822179Z","shell.execute_reply":"2025-04-04T13:58:10.145569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"taxonomy_df = pd.read_csv(cfg.taxonomy_csv)\ntaxonomy_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.147424Z","iopub.execute_input":"2025-04-04T13:58:10.147695Z","iopub.status.idle":"2025-04-04T13:58:10.180477Z","shell.execute_reply.started":"2025-04-04T13:58:10.147675Z","shell.execute_reply":"2025-04-04T13:58:10.179706Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"taxonomy_df.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.181422Z","iopub.execute_input":"2025-04-04T13:58:10.181736Z","iopub.status.idle":"2025-04-04T13:58:10.206434Z","shell.execute_reply.started":"2025-04-04T13:58:10.181706Z","shell.execute_reply":"2025-04-04T13:58:10.205479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"unique_primary_labels = taxonomy_df[\"primary_label\"].unique()\nspecies_ids = dict()\nint_code = 0\nfor elt in unique_primary_labels:\n    species_ids[elt] = int_code\n    int_code +=1\n\n# species_ids","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.208923Z","iopub.execute_input":"2025-04-04T13:58:10.209266Z","iopub.status.idle":"2025-04-04T13:58:10.224319Z","shell.execute_reply.started":"2025-04-04T13:58:10.209232Z","shell.execute_reply":"2025-04-04T13:58:10.223387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(cfg.train_csv)\ntrain_df.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.225495Z","iopub.execute_input":"2025-04-04T13:58:10.225740Z","iopub.status.idle":"2025-04-04T13:58:10.425185Z","shell.execute_reply.started":"2025-04-04T13:58:10.225719Z","shell.execute_reply":"2025-04-04T13:58:10.424226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_rows = train_df.shape[0]\nnum_rows_train = math.ceil(num_rows * cfg.train_perc)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.426233Z","iopub.execute_input":"2025-04-04T13:58:10.426548Z","iopub.status.idle":"2025-04-04T13:58:10.431638Z","shell.execute_reply.started":"2025-04-04T13:58:10.426519Z","shell.execute_reply":"2025-04-04T13:58:10.430703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df[\"primary_label\"].unique()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.432557Z","iopub.execute_input":"2025-04-04T13:58:10.432769Z","iopub.status.idle":"2025-04-04T13:58:10.453762Z","shell.execute_reply.started":"2025-04-04T13:58:10.432751Z","shell.execute_reply":"2025-04-04T13:58:10.452840Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdCLEFDataset(Dataset):\n    def __init__(self, df, data_path):\n        self.df = df\n        self.data_path = data_path\n        self.cfg = cfg\n        self.target_length = 1024  # AST typically expects this length\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        audio_path = os.path.join(self.data_path, row['filename'])\n        \n        # Load audio and ensure it's the right length\n        waveform, _ = librosa.load(audio_path, sr=16000, mono=True)\n        \n        # Pad or truncate to target length\n        if len(waveform) > self.target_length:\n            waveform = waveform[:self.target_length]\n        else:\n            padding = self.target_length - len(waveform)\n            waveform = np.pad(waveform, (0, padding), mode='constant')\n        \n        # Convert to tensor and add batch dimension\n        waveform = torch.tensor(waveform).float()\n        waveform = waveform.squeeze().numpy()\n        \n        label = species_ids[row['primary_label']]\n        return {\n            'input_values': waveform,\n            'label': label\n        }\n        # return waveform, torch.tensor(label, dtype=torch.long)\n    \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.454555Z","iopub.execute_input":"2025-04-04T13:58:10.454885Z","iopub.status.idle":"2025-04-04T13:58:10.472219Z","shell.execute_reply.started":"2025-04-04T13:58:10.454850Z","shell.execute_reply":"2025-04-04T13:58:10.471466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_fn(batch):\n    audio_feats = [x['input_values'] for x in batch]\n    labels = [x['label'] for x in batch]\n    labels = torch.tensor(labels)\n    out_aud_feats = feature_extractor(audio_feats,\n                                      sampling_rate=16000,\n                                      return_tensors='pt')['input_values']\n    out_aud_feats = out_aud_feats.half()\n    return out_aud_feats, labels\n\n# def collate_fn(batch):\n#     waveforms, labels = zip(*batch)\n    \n#     # Stack waveforms and add channel dimension (AST expects [batch_size, 1, sequence_length])\n#     waveforms = torch.stack(waveforms).unsqueeze(1)\n#     labels = torch.stack(labels)\n    \n#     return waveforms, labels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.472984Z","iopub.execute_input":"2025-04-04T13:58:10.473261Z","iopub.status.idle":"2025-04-04T13:58:10.487279Z","shell.execute_reply.started":"2025-04-04T13:58:10.473232Z","shell.execute_reply":"2025-04-04T13:58:10.486449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_model(num_classes):\n    model = ASTForAudioClassification.from_pretrained(cfg.model_name, num_labels=num_classes,\n                                                      ignore_mismatched_sizes=True,\n                                                     torch_dtype=torch.float16)\n    return model.to(cfg.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.488094Z","iopub.execute_input":"2025-04-04T13:58:10.488404Z","iopub.status.idle":"2025-04-04T13:58:10.506059Z","shell.execute_reply.started":"2025-04-04T13:58:10.488382Z","shell.execute_reply":"2025-04-04T13:58:10.505214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, train_loader, optimizer, scheduler, criterion):\n    model.train()\n    total_loss = 0.0\n    for batch in tqdm(train_loader):\n        # print(\"Input shape:\", batch[0].shape)  # Devrait être (batch_size, 1, time, frequency)\n        # break\n        inputs, labels = batch\n        inputs, labels = inputs.to(cfg.device), labels.to(cfg.device)\n        \n        optimizer.zero_grad()\n        outputs = model(inputs).logits\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n        total_loss += loss.item()\n        break\n    \n    return total_loss / len(train_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.506884Z","iopub.execute_input":"2025-04-04T13:58:10.507089Z","iopub.status.idle":"2025-04-04T13:58:10.522091Z","shell.execute_reply.started":"2025-04-04T13:58:10.507071Z","shell.execute_reply":"2025-04-04T13:58:10.521249Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, loader, criterion):\n    val_loss, val_correct, val_total = 0, 0, 0\n    model.eval()\n    losses = []\n    all_targets = []\n    all_outputs = []\n\n    with torch.no_grad():\n        for batch in tqdm(loader, desc=\"Validating\"):\n            inputs, labels = batch\n            inputs, labels = inputs.to(cfg.device), labels.to(cfg.device)\n            outputs = model(inputs).logits\n            loss = criterion(outputs, labels)\n            val_loss += loss.item()\n            preds = torch.argmax(outputs, dim=1)\n            val_correct += (preds == labels).sum().item()\n            val_total += labels.size(0)\n            break\n    # Log validation metrics\n    val_acc = (val_correct / val_total) * 100\n    print(f\"Validation Loss: {val_loss / len(loader):.4f}, Validation Accuracy: {val_acc:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.522886Z","iopub.execute_input":"2025-04-04T13:58:10.523131Z","iopub.status.idle":"2025-04-04T13:58:10.539107Z","shell.execute_reply.started":"2025-04-04T13:58:10.523112Z","shell.execute_reply":"2025-04-04T13:58:10.538359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train():\n    df = pd.read_csv(cfg.train_csv)\n    df = df.sample(frac=1.).reset_index(drop=True) # we shuffle the whole table\n    df_train = df.head(num_rows_train)\n    num_rows_val = num_rows - num_rows_train\n    df_val = df.tail(num_rows_val)\n    df_val = df_val.reset_index(drop=True)\n    \n    dataset_train = BirdCLEFDataset(df_train, cfg.train_datadir)\n    dataset_val = BirdCLEFDataset(df_val, cfg.train_datadir)\n    train_loader = DataLoader(dataset_train,\n                              batch_size=cfg.batch_size,\n                              shuffle=True,\n                              #num_workers=cfg.num_workers,\n                              collate_fn=collate_fn)\n    val_loader = DataLoader(dataset_val,\n                              batch_size=cfg.batch_size,\n                              shuffle=True,\n                              #num_workers=cfg.num_workers,\n                              collate_fn=collate_fn)\n\n    \n    \n    taxonomy_df = pd.read_csv(cfg.taxonomy_csv)\n    species_ids = taxonomy_df['primary_label'].tolist()\n    # self.num_classes = len(self.species_ids)\n    \n    model = get_model(num_classes=len(species_ids))\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n    scheduler = get_scheduler(\"linear\", optimizer=optimizer, num_warmup_steps=500, num_training_steps=len(train_loader) * cfg.epochs)\n    \n    for epoch in range(cfg.epochs):\n        loss = train_one_epoch(model, train_loader, optimizer, scheduler, criterion)\n        print(f\"Epoch {epoch+1}/{cfg.epochs}, Loss: {loss:.4f}\")\n        validate(model, val_loader, criterion)\n    \n    # torch.save(model.state_dict(), os.path.join(cfg.OUTPUT_DIR, \"birdclef_ast.pth\"))\n        model.save_pretrained(f\"fine_tuned_ast_epoch_{epoch}\")\n    print(\"Model saved!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.539823Z","iopub.execute_input":"2025-04-04T13:58:10.540033Z","iopub.status.idle":"2025-04-04T13:58:10.559858Z","shell.execute_reply.started":"2025-04-04T13:58:10.540014Z","shell.execute_reply":"2025-04-04T13:58:10.558934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    train()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T13:58:10.560693Z","iopub.execute_input":"2025-04-04T13:58:10.560888Z","iopub.status.idle":"2025-04-04T13:58:42.564897Z","shell.execute_reply.started":"2025-04-04T13:58:10.560871Z","shell.execute_reply":"2025-04-04T13:58:42.564146Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}