{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# -----------------------------\n# 1. BUSINESS UNDERSTANDING 📈\n# -----------------------------\n## Goal: Classify bird species from audio recordings in real environments.\n## Evaluation Metric: Macro-averaged ROC-AUC, skipping species with no true positives.\n### https://www.kaggle.com/code/metric/birdclef-roc-auc","metadata":{}},{"cell_type":"markdown","source":"# -----------------------------\n# 2. DATA UNDERSTANDING 🔍\n# -----------------------------","metadata":{}},{"cell_type":"code","source":"!pip install mlflow","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:10:52.682734Z","iopub.execute_input":"2025-04-14T13:10:52.683047Z","iopub.status.idle":"2025-04-14T13:11:03.024611Z","shell.execute_reply.started":"2025-04-14T13:10:52.683009Z","shell.execute_reply":"2025-04-14T13:11:03.023831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install mutagen\n!pip install contextily\n!pip install silero_vad","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:11:07.713644Z","iopub.execute_input":"2025-04-14T13:11:07.714008Z","iopub.status.idle":"2025-04-14T13:11:22.439872Z","shell.execute_reply.started":"2025-04-14T13:11:07.713976Z","shell.execute_reply":"2025-04-14T13:11:22.439045Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install sweetviz","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:11:26.658418Z","iopub.execute_input":"2025-04-14T13:11:26.658702Z","iopub.status.idle":"2025-04-14T13:11:30.715300Z","shell.execute_reply.started":"2025-04-14T13:11:26.658678Z","shell.execute_reply":"2025-04-14T13:11:30.714115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"######## Common ########\nimport pandas as pd\nimport os\nimport copy\nimport numpy as np\nimport copy\nfrom collections import defaultdict, Counter\n\n######## Data Processing ########\nfrom sklearn.model_selection import train_test_split, GridSearchCV\nfrom sklearn.metrics import accuracy_score, classification_report, confusion_matrix\nfrom sklearn.preprocessing import LabelEncoder\n\n######## Sound ########\nimport librosa\nfrom mutagen.oggvorbis import OggVorbis\n\n# Human Voices Detection\nfrom silero_vad import get_speech_timestamps, save_audio\n\n######## Machine Learning ########\nimport torch\nimport torchaudio\nimport xgboost as xgb\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.callbacks import ModelCheckpoint\n######## Visualization ########\n# Ploting\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n# Geographical Maps\nimport geopandas as gpd\nimport contextily as ctx\n\n######## Logging ########\n\nimport mlflow\nimport mlflow.sklearn\nimport tempfile\n\n\n######## EDA ########\nimport sweetviz as sv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:11:33.676773Z","iopub.execute_input":"2025-04-14T13:11:33.677106Z","iopub.status.idle":"2025-04-14T13:11:47.396608Z","shell.execute_reply.started":"2025-04-14T13:11:33.677080Z","shell.execute_reply":"2025-04-14T13:11:47.395974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n\n# Load CSV and TXT metadata\ndata_df = pd.read_csv('/kaggle/input/birdclef-2025/train.csv')\ntaxonomy_df = pd.read_csv('/kaggle/input/birdclef-2025/taxonomy.csv')\nsubmission_df = pd.read_csv('/kaggle/input/birdclef-2025/sample_submission.csv')\nlocation_df = pd.read_csv('/kaggle/input/birdclef-2025/recording_location.txt', delimiter='\\t')\n\n# Quick overview\nprint(\"Train shape:\", data_df.shape)\nprint(\"Taxonomy shape:\", taxonomy_df.shape)\ndata_df.head()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:11:54.739062Z","iopub.execute_input":"2025-04-14T13:11:54.739608Z","iopub.status.idle":"2025-04-14T13:11:54.962426Z","shell.execute_reply.started":"2025-04-14T13:11:54.739580Z","shell.execute_reply":"2025-04-14T13:11:54.961508Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### EDA:","metadata":{}},{"cell_type":"code","source":"# Generate SweetViz report for data_df\nreport_data = sv.analyze(data_df)\nreport_data.show_html('SweetViz_report_data_df.html')  # This will generate an HTML file\n\n# Generate SweetViz report for taxonomy_df\nreport_taxonomy = sv.analyze(taxonomy_df)\nreport_taxonomy.show_html('SweetViz_report_taxonomy_df.html')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:04.702876Z","iopub.execute_input":"2025-04-14T13:12:04.703187Z","iopub.status.idle":"2025-04-14T13:12:09.168052Z","shell.execute_reply.started":"2025-04-14T13:12:04.703160Z","shell.execute_reply":"2025-04-14T13:12:09.167378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"report_data.show_html('SweetViz_report_data_df.html')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:13.959373Z","iopub.execute_input":"2025-04-14T13:12:13.959657Z","iopub.status.idle":"2025-04-14T13:12:13.982042Z","shell.execute_reply.started":"2025-04-14T13:12:13.959636Z","shell.execute_reply":"2025-04-14T13:12:13.981189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_df = pd.read_csv('/kaggle/input/birdclef-2025/train.csv')\n# Label Mapping\nlabels = data_df['primary_label'].unique()\n# Sort for consistency\nsorted_labels = sorted(labels)\n\n# Creating dictionaries\nlabel_to_common = taxonomy_df.set_index(\"primary_label\")[\"common_name\"].to_dict()\nlabel_to_class = taxonomy_df.set_index(\"primary_label\")[\"class_name\"].to_dict()\n\nprint(\"Shape:\", data_df.shape)\ndata_df.info()\ndata_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:16.010892Z","iopub.execute_input":"2025-04-14T13:12:16.011207Z","iopub.status.idle":"2025-04-14T13:12:16.157078Z","shell.execute_reply.started":"2025-04-14T13:12:16.011179Z","shell.execute_reply":"2025-04-14T13:12:16.156238Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Filtering\n### Remove uneccessry columns","metadata":{}},{"cell_type":"code","source":"data_df.drop(columns=['author', 'license', 'url', 'latitude', 'longitude', 'type'],inplace=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:27.154318Z","iopub.execute_input":"2025-04-14T13:12:27.154670Z","iopub.status.idle":"2025-04-14T13:12:27.162164Z","shell.execute_reply.started":"2025-04-14T13:12:27.154631Z","shell.execute_reply":"2025-04-14T13:12:27.161284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_df.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:30.856691Z","iopub.execute_input":"2025-04-14T13:12:30.857018Z","iopub.status.idle":"2025-04-14T13:12:30.873164Z","shell.execute_reply.started":"2025-04-14T13:12:30.856993Z","shell.execute_reply":"2025-04-14T13:12:30.872343Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Removing Samples with secondary labels (examples where multiple animals make sound) due insufficiant data for multilabeling","metadata":{}},{"cell_type":"code","source":"# Remove data with secondary labels\ndata_df = data_df[data_df['secondary_labels'] == \"['']\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:34.271662Z","iopub.execute_input":"2025-04-14T13:12:34.272020Z","iopub.status.idle":"2025-04-14T13:12:34.281786Z","shell.execute_reply.started":"2025-04-14T13:12:34.271991Z","shell.execute_reply":"2025-04-14T13:12:34.280874Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Add animal common name and class columns","metadata":{}},{"cell_type":"code","source":"# Apply dictionaries to create new columns\ndata_df[\"common_name\"] = data_df[\"primary_label\"].map(label_to_common)\ndata_df[\"animal_class\"] = data_df[\"primary_label\"].map(label_to_class)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:36.867926Z","iopub.execute_input":"2025-04-14T13:12:36.868224Z","iopub.status.idle":"2025-04-14T13:12:36.877537Z","shell.execute_reply.started":"2025-04-14T13:12:36.868203Z","shell.execute_reply":"2025-04-14T13:12:36.876734Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:39.737765Z","iopub.execute_input":"2025-04-14T13:12:39.738077Z","iopub.status.idle":"2025-04-14T13:12:39.750577Z","shell.execute_reply.started":"2025-04-14T13:12:39.738056Z","shell.execute_reply":"2025-04-14T13:12:39.749853Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"report_data = sv.analyze(data_df)\nreport_data.show_html('SweetViz_report_data_df.html')  # This will generate an HTML file","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:44.297323Z","iopub.execute_input":"2025-04-14T13:12:44.297596Z","iopub.status.idle":"2025-04-14T13:12:45.712946Z","shell.execute_reply.started":"2025-04-14T13:12:44.297573Z","shell.execute_reply":"2025-04-14T13:12:45.712268Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Filter dataframe to contain only select spicies from each class","metadata":{}},{"cell_type":"code","source":"num_of_top_most_common = 4\ninsecta_species = data_df[data_df[\"animal_class\"] == \"Insecta\"][\"common_name\"].value_counts().nlargest(num_of_top_most_common).index\nmammalia_species = data_df[data_df[\"animal_class\"] == \"Mammalia\"][\"common_name\"].value_counts().nlargest(num_of_top_most_common).index\namphibia_species = data_df[data_df[\"animal_class\"] == \"Amphibia\"][\"common_name\"].value_counts().nlargest(num_of_top_most_common).index\naves_species = data_df[data_df[\"animal_class\"] == \"Aves\"][\"common_name\"].value_counts().nlargest(num_of_top_most_common).index\n\n# Filter train_df to include only the top species from each animal class\nfiltered_df = data_df[\n    (data_df[\"common_name\"].isin(insecta_species)) |\n    (data_df[\"common_name\"].isin(mammalia_species)) |\n    (data_df[\"common_name\"].isin(amphibia_species)) |\n    (data_df[\"common_name\"].isin(aves_species))\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:50.539305Z","iopub.execute_input":"2025-04-14T13:12:50.539592Z","iopub.status.idle":"2025-04-14T13:12:50.566492Z","shell.execute_reply.started":"2025-04-14T13:12:50.539569Z","shell.execute_reply":"2025-04-14T13:12:50.565870Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Downsample data","metadata":{}},{"cell_type":"code","source":"max_samples_per_animal = 33\n\n# Downsample each species to at most 33 samples\ndownsampled_df = filtered_df.groupby(\"common_name\", group_keys=False).apply(\n    lambda x: x.sample(n=min(len(x), max_samples_per_animal), random_state=42)\n).reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:54.376029Z","iopub.execute_input":"2025-04-14T13:12:54.376308Z","iopub.status.idle":"2025-04-14T13:12:54.400086Z","shell.execute_reply.started":"2025-04-14T13:12:54.376285Z","shell.execute_reply":"2025-04-14T13:12:54.399136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compare original and cleaned\ncompare_report = sv.compare([data_df, \"Original\"], [downsampled_df, \"Downsampled\"])\ncompare_report.show_html(\"SweetViz_comparison.html\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:12:57.206517Z","iopub.execute_input":"2025-04-14T13:12:57.206855Z","iopub.status.idle":"2025-04-14T13:12:59.398893Z","shell.execute_reply.started":"2025-04-14T13:12:57.206825Z","shell.execute_reply":"2025-04-14T13:12:59.398209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Check the balance of classes and animals","metadata":{}},{"cell_type":"code","source":"def show_class_and_animal_balance(train_df):\n    \n    # Create dataframe with animal name, animal class, and count\n    animal_counts_df = train_df.groupby([\"common_name\", \"animal_class\"]).size().reset_index(name=\"count\")\n    \n    for c in [\"Aves\",\"Mammalia\",\"Amphibia\",\"Insecta\"]:\n        df = animal_counts_df[animal_counts_df['animal_class']==c]\n    \n        df.sort_values(by=\"count\", ascending=True, inplace=True)\n        \n        # Plot horizontal bar chart\n        plt.figure(figsize=(3, 2))  # Adjust height for better visualization\n        ax = sns.barplot(x=\"count\", y=\"common_name\", data=df, hue=\"common_name\", palette=\"viridis\")\n        if ax.legend_:\n            ax.legend_.remove()\n        \n        \n        # Add count numbers to each bar\n        for container in ax.containers:\n            ax.bar_label(container, fmt='%d', padding=5)\n        \n        # Adjust x-axis limit for spacing\n        max_count = df[\"count\"].max()\n        plt.xlim(0, max_count * 1.2)  # Extend the max value by 15%\n        \n        plt.xlabel(\"Number of Recordings\")\n        plt.ylabel(\"common_name\")\n        plt.title(f\"Number of Recordings per {c} Animal\")\n        plt.grid(axis=\"x\", linestyle=\"--\", alpha=0.6)\n        plt.show()\n\nshow_class_and_animal_balance(downsampled_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:13:02.102435Z","iopub.execute_input":"2025-04-14T13:13:02.102719Z","iopub.status.idle":"2025-04-14T13:13:02.864570Z","shell.execute_reply.started":"2025-04-14T13:13:02.102696Z","shell.execute_reply":"2025-04-14T13:13:02.863632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Stratified train/validation/test split\ndef stratified_split(df, stratify_col, train_size=0.8, val_size=0.1, test_size=0.1):\n    train_df, temp_df = train_test_split(df, train_size=train_size, stratify=df[stratify_col], random_state=42)\n    val_df, test_df = train_test_split(temp_df, train_size=val_size / (val_size + test_size), stratify=temp_df[stratify_col], random_state=42)\n    return train_df, val_df, test_df\n\ntrain_df, val_df, test_df = stratified_split(downsampled_df, stratify_col=[\"animal_class\", \"primary_label\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:13:07.039493Z","iopub.execute_input":"2025-04-14T13:13:07.039837Z","iopub.status.idle":"2025-04-14T13:13:07.055035Z","shell.execute_reply.started":"2025-04-14T13:13:07.039807Z","shell.execute_reply":"2025-04-14T13:13:07.054163Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:13:09.431596Z","iopub.execute_input":"2025-04-14T13:13:09.431961Z","iopub.status.idle":"2025-04-14T13:13:09.444742Z","shell.execute_reply.started":"2025-04-14T13:13:09.431919Z","shell.execute_reply":"2025-04-14T13:13:09.443999Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"show_class_and_animal_balance(test_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:13:12.152665Z","iopub.execute_input":"2025-04-14T13:13:12.153006Z","iopub.status.idle":"2025-04-14T13:13:12.941912Z","shell.execute_reply.started":"2025-04-14T13:13:12.152976Z","shell.execute_reply":"2025-04-14T13:13:12.941197Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compare original and cleaned\ncompare_report = sv.compare([train_df, \"Train\"], [test_df, \"Test\"])\ncompare_report.show_html(\"Train_test_split.html\")\ncompare_report_val = sv.compare([train_df, \"Train\"], [val_df, \"Validation\"])\ncompare_report_val.show_html(\"Train_val_split.html\")\ncompare_report_val_test = sv.compare([test_df, \"Test\"], [val_df, \"Validation\"])\ncompare_report_val_test.show_html(\"Test_val_split.html\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:13:16.391530Z","iopub.execute_input":"2025-04-14T13:13:16.391844Z","iopub.status.idle":"2025-04-14T13:13:22.810427Z","shell.execute_reply.started":"2025-04-14T13:13:16.391816Z","shell.execute_reply":"2025-04-14T13:13:22.809759Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Human Speech Recognition in Recordings","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nimport torchaudio\n\ndef seconds_to_mmss(seconds):\n    mins = int(seconds // 60)\n    secs = int(seconds % 60)\n    return f\"{mins:02d}:{secs:02d}\"\n\ndef extract_speech_timestamps_for_df(df, audio_root, threshold=0.17):\n    model, utils = torch.hub.load('snakers4/silero-vad', 'silero_vad', force_reload=False, trust_repo=True)\n    get_speech_timestamps, *_ = utils\n    \n    results = []\n\n    for _, row in tqdm(df.iterrows(), total=len(df)):\n        filename = row['filename']\n        file_path = os.path.join(audio_root, filename)\n\n        try:\n            wav, sr = torchaudio.load(file_path)\n            raw_timestamps = get_speech_timestamps(wav, model, sampling_rate=sr, threshold=threshold)\n\n            # Convert samples to readable time (mm:ss)\n            timestamps_mmss = []\n            for segment in raw_timestamps:\n                start_sec = segment['start'] / sr\n                end_sec = segment['end'] / sr\n                timestamps_mmss.append({\n                    'start': seconds_to_mmss(start_sec),\n                    'end': seconds_to_mmss(end_sec)\n                })\n\n            result = {\n                'filename': filename,\n                'speech_timestamps': timestamps_mmss\n            }\n\n            # Add all original metadata\n            for col in df.columns:\n                result[col] = row[col]\n\n            results.append(result)\n\n        except Exception as e:\n            print(f\"[ERROR] Could not process {filename}: {e}\")\n            continue\n\n    return results\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:13:30.732286Z","iopub.execute_input":"2025-04-14T13:13:30.732569Z","iopub.status.idle":"2025-04-14T13:13:30.739249Z","shell.execute_reply.started":"2025-04-14T13:13:30.732546Z","shell.execute_reply":"2025-04-14T13:13:30.738526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the audio path on your local/Kaggle environment\naudio_path = \"/kaggle/input/birdclef-2025/train_audio\"\n\n# Extract timestamps for training set\nspeech_info_train = extract_speech_timestamps_for_df(train_df, audio_root=audio_path)\n\n# Example output per item:\n# {\n#     'filename': '48124/CSA35116.ogg',\n#     'speech_timestamps': [{'start': 1000, 'end': 32000}, ...],\n#     'primary_label': 'compau',\n#     ...\n# }\n\n# Convert to DataFrame if needed\nspeech_df_train = pd.DataFrame(speech_info_train)\nspeech_df_train.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:13:33.670035Z","iopub.execute_input":"2025-04-14T13:13:33.670313Z","iopub.status.idle":"2025-04-14T13:19:21.948983Z","shell.execute_reply.started":"2025-04-14T13:13:33.670290Z","shell.execute_reply":"2025-04-14T13:19:21.948240Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract timestamps for training set\nspeech_info_val = extract_speech_timestamps_for_df(val_df, audio_root=audio_path)\n\n# Example output per item:\n# {\n#     'filename': '48124/CSA35116.ogg',\n#     'speech_timestamps': [{'start': 1000, 'end': 32000}, ...],\n#     'primary_label': 'compau',\n#     ...\n# }\n\n# Convert to DataFrame if needed\nspeech_df_val = pd.DataFrame(speech_info_val)\nspeech_df_val.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:19:21.950044Z","iopub.execute_input":"2025-04-14T13:19:21.950285Z","iopub.status.idle":"2025-04-14T13:20:02.911075Z","shell.execute_reply.started":"2025-04-14T13:19:21.950264Z","shell.execute_reply":"2025-04-14T13:20:02.910159Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract timestamps for training set\nspeech_info_test = extract_speech_timestamps_for_df(test_df, audio_root=audio_path)\n\n# Example output per item:\n# {\n#     'filename': '48124/CSA35116.ogg',\n#     'speech_timestamps': [{'start': 1000, 'end': 32000}, ...],\n#     'primary_label': 'compau',\n#     ...\n# }\n\n# Convert to DataFrame if needed\nspeech_df_test = pd.DataFrame(speech_info_test)\nspeech_df_test.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:20:02.912558Z","iopub.execute_input":"2025-04-14T13:20:02.912836Z","iopub.status.idle":"2025-04-14T13:20:45.082973Z","shell.execute_reply.started":"2025-04-14T13:20:02.912782Z","shell.execute_reply":"2025-04-14T13:20:45.082147Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Extracting audio features for all files in each df:","metadata":{}},{"cell_type":"code","source":"import os\nimport librosa\nimport pandas as pd\nimport numpy as np\nfrom tqdm import tqdm\n\ndef normalize_feature(X):\n    return (X - X.min()) / (X.max() - X.min() + 1e-8)\n\ndef mmss_to_seconds(mmss: str):\n    \"\"\"Convert 'mm:ss' format to total seconds\"\"\"\n    mins, secs = map(int, mmss.split(\":\"))\n    return mins * 60 + secs\n\ndef extract_features_for_df_with_speech_removal(df, audio_root, chunk_duration=5, speech_df=None):\n    all_rows = []\n\n    # Convert speech_df to dictionary: filename → list of timestamps\n    speech_map = {}\n    if speech_df is not None:\n        for row in speech_df:\n            speech_map[row['filename']] = row['speech_timestamps']\n\n    for i, row in tqdm(df.iterrows(), total=len(df)):\n        filename = row['filename']\n        file_path = os.path.join(audio_root, filename)\n\n        try:\n            y, sr = librosa.load(file_path, sr=None)\n        except Exception as e:\n            print(f\"Error loading {file_path}: {e}\")\n            continue\n\n        duration = len(y) / sr  # in seconds\n\n        # === Remove speech if present ===\n        if filename in speech_map:\n            speech = speech_map[filename]\n            if speech:  # Check non-empty list\n                speech_starts = [mmss_to_seconds(seg['start']) for seg in speech]\n                speech_ends = [mmss_to_seconds(seg['end']) for seg in speech]\n\n                earliest = min(speech_starts)\n                latest = max(speech_ends)\n\n                if latest < duration / 2:\n                    # Speech is at the beginning → remove from start to ceil(latest/5)*5\n                    trim_sec = int(np.ceil(latest / chunk_duration) * chunk_duration)\n                    y = y[int(trim_sec * sr):]\n                elif earliest > duration / 2:\n                    # Speech is at the end → remove from floor(earliest/5)*5 to end\n                    trim_sec = int(np.floor(earliest / chunk_duration) * chunk_duration)\n                    y = y[:int(trim_sec * sr)]\n                else:\n                    pass\n\n        samples_per_chunk = int(sr * chunk_duration)\n        n_chunks = int(len(y) / samples_per_chunk)\n\n        if n_chunks == 0:\n            pass\n\n\n        for c in range(n_chunks):\n            start = c * samples_per_chunk\n            end = start + samples_per_chunk\n            y_chunk = y[start:end]\n\n            # === Extract Features ===\n            mfcc = librosa.feature.mfcc(y=y_chunk, sr=sr, n_mfcc=20)\n            chroma = librosa.feature.chroma_stft(y=y_chunk, sr=sr)\n            contrast = librosa.feature.spectral_contrast(y=y_chunk, sr=sr)\n            zcr = librosa.feature.zero_crossing_rate(y_chunk)\n            rms = librosa.feature.rms(y=y_chunk)\n            centroid = librosa.feature.spectral_centroid(y=y_chunk, sr=sr)\n            bandwidth = librosa.feature.spectral_bandwidth(y=y_chunk, sr=sr)\n            rolloff = librosa.feature.spectral_rolloff(y=y_chunk, sr=sr)\n\n            # === Normalize ===\n            mfcc = normalize_feature(mfcc)\n            chroma = normalize_feature(chroma)\n            contrast = normalize_feature(contrast)\n            zcr = normalize_feature(zcr)\n            rms = normalize_feature(rms)\n            centroid = normalize_feature(centroid)\n            bandwidth = normalize_feature(bandwidth)\n            rolloff = normalize_feature(rolloff)\n\n            # === Average features ===\n            chunk_features = {\n                f'mfcc_{i}': mfcc[i].mean() for i in range(mfcc.shape[0])\n            }\n            chunk_features.update({\n                f'chroma_{i}': chroma[i].mean() for i in range(chroma.shape[0])\n            })\n            chunk_features.update({\n                f'contrast_{i}': contrast[i].mean() for i in range(contrast.shape[0])\n            })\n            chunk_features.update({\n                'zcr': zcr.mean(),\n                'rms': rms.mean(),\n                'centroid': centroid.mean(),\n                'bandwidth': bandwidth.mean(),\n                'rolloff': rolloff.mean(),\n                'chunk_start_time': c * chunk_duration  # Add timestamp\n            })\n\n            for col in df.columns:\n                chunk_features[col] = row[col]\n\n            all_rows.append(chunk_features)\n\n    return pd.DataFrame(all_rows)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:22:10.223314Z","iopub.execute_input":"2025-04-14T13:22:10.223645Z","iopub.status.idle":"2025-04-14T13:22:10.236989Z","shell.execute_reply.started":"2025-04-14T13:22:10.223615Z","shell.execute_reply":"2025-04-14T13:22:10.235997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_features = extract_features_for_df_with_speech_removal(train_df, audio_root='/kaggle/input/birdclef-2025/train_audio', speech_df=speech_info_train)\nval_features   = extract_features_for_df_with_speech_removal(val_df, audio_root='/kaggle/input/birdclef-2025/train_audio', speech_df=speech_info_val)\ntest_features  = extract_features_for_df_with_speech_removal(test_df, audio_root='/kaggle/input/birdclef-2025/train_audio', speech_df=speech_info_test)\n\n\ntrain_features.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:22:55.090240Z","iopub.execute_input":"2025-04-14T13:22:55.090559Z","iopub.status.idle":"2025-04-14T13:31:53.188504Z","shell.execute_reply.started":"2025-04-14T13:22:55.090530Z","shell.execute_reply":"2025-04-14T13:31:53.187809Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_features","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:31:57.455668Z","iopub.execute_input":"2025-04-14T13:31:57.456022Z","iopub.status.idle":"2025-04-14T13:31:57.481132Z","shell.execute_reply.started":"2025-04-14T13:31:57.455993Z","shell.execute_reply":"2025-04-14T13:31:57.480214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"speech_df_train.to_csv(\"speech_df_train.csv\", index=False)\nspeech_df_val.to_csv(\"speech_df_val.csv\", index=False)\nspeech_df_test.to_csv(\"speech_df_test.csv\", index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:32:01.217436Z","iopub.execute_input":"2025-04-14T13:32:01.217726Z","iopub.status.idle":"2025-04-14T13:32:01.233482Z","shell.execute_reply.started":"2025-04-14T13:32:01.217695Z","shell.execute_reply":"2025-04-14T13:32:01.232690Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"taxonomy_df = pd.read_csv('/kaggle/input/birdclef-2025/taxonomy.csv')\nprint(\"Shape:\", taxonomy_df.shape)\ntaxonomy_df.info()\ntaxonomy_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:32:03.700326Z","iopub.execute_input":"2025-04-14T13:32:03.700637Z","iopub.status.idle":"2025-04-14T13:32:03.721867Z","shell.execute_reply.started":"2025-04-14T13:32:03.700607Z","shell.execute_reply":"2025-04-14T13:32:03.721205Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Baseline: XGBoost Classifier","metadata":{}},{"cell_type":"code","source":"def prepare_xgb_data(train_df, val_df):\n    drop_cols = ['filename', 'primary_label', 'common_name', 'scientific_name', 'collection', 'secondary_labels', 'rating']\n    features = [col for col in train_df.columns if col not in drop_cols + ['animal_class']]\n    \n    le = LabelEncoder()\n    train_labels = le.fit_transform(train_df['animal_class'])\n    val_labels = le.transform(val_df['animal_class'])\n    \n    train_X = train_df[features]\n    val_X = val_df[features]\n    \n    return train_X, train_labels, val_X, val_labels, le\n\ntrain_X, train_y, val_X, val_y, label_encoder = prepare_xgb_data(train_features, val_features)\n\n# MLflow Experiment Logging\nwith mlflow.start_run(run_name=\"XGBoost_BirdCLEF\"):\n    # Log parameters\n    mlflow.log_param(\"num_class\", 4)\n    mlflow.log_param(\"n_estimators\", 100)\n    mlflow.log_param(\"max_depth\", 6)\n    mlflow.log_param(\"learning_rate\", 0.1)\n\n    # Train model\n    model = xgb.XGBClassifier(\n        objective='multi:softprob',\n        num_class=4,\n        eval_metric='mlogloss',\n        use_label_encoder=False,\n        n_estimators=100,\n        max_depth=6,\n        learning_rate=0.1\n    )\n    model.fit(train_X, train_y)\n\n    # Validation evaluation\n    val_preds = model.predict(val_X)\n    val_acc = accuracy_score(val_y, val_preds)\n\n    print(\"Accuracy:\", val_acc)\n    print(\"Classification Report:\\n\", classification_report(val_y, val_preds, target_names=label_encoder.classes_))\n\n    # Log accuracy\n    mlflow.log_metric(\"val_accuracy\", val_acc)\n\n    # Log confusion matrix\n    cm = confusion_matrix(val_y, val_preds)\n    plt.figure(figsize=(6, 5))\n    sns.heatmap(cm, annot=True, fmt=\"d\", xticklabels=label_encoder.classes_, yticklabels=label_encoder.classes_, cmap=\"Blues\")\n    plt.title(\"Validation Confusion Matrix\")\n    plt.xlabel(\"Predicted\")\n    plt.ylabel(\"True\")\n    plt.tight_layout()\n\n    with tempfile.NamedTemporaryFile(suffix=\".png\", delete=False) as tmp:\n        plt.savefig(tmp.name)\n        mlflow.log_artifact(tmp.name, artifact_path=\"val_confusion_matrix\")\n\n    # Log model\n    mlflow.sklearn.log_model(model, \"xgb_model\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-10T12:46:59.942272Z","iopub.execute_input":"2025-04-10T12:46:59.942480Z","iopub.status.idle":"2025-04-10T12:47:14.917354Z","shell.execute_reply.started":"2025-04-10T12:46:59.942462Z","shell.execute_reply":"2025-04-10T12:47:14.916424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_test_data(test_df, label_encoder, features):\n    test_X = test_df[features]\n    test_y = label_encoder.transform(test_df['animal_class'])\n    return test_X, test_y\n\nused_features = train_X.columns.tolist()\ntest_X, test_y = prepare_test_data(test_features, label_encoder, used_features)\ntest_preds = model.predict(test_X)\n\ntest_acc = accuracy_score(test_y, test_preds)\nprint(\"✅ Test Set Accuracy:\", test_acc)\nprint(\"✅ Test Set Classification Report:\\n\", classification_report(test_y, test_preds, target_names=label_encoder.classes_))\nmlflow.log_metric(\"test_accuracy\", test_acc)\n\n# Test confusion matrix\ncm = confusion_matrix(test_y, test_preds)\nplt.figure(figsize=(6, 5))\nsns.heatmap(cm, annot=True, fmt=\"d\", xticklabels=label_encoder.classes_, yticklabels=label_encoder.classes_, cmap=\"Blues\")\nplt.title(\"Test Set Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.tight_layout()\nplt.show()\n\nwith tempfile.NamedTemporaryFile(suffix=\".png\", delete=False) as tmp:\n    plt.savefig(tmp.name)\n    mlflow.log_artifact(tmp.name, artifact_path=\"test_confusion_matrix\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-10T12:47:14.918459Z","iopub.execute_input":"2025-04-10T12:47:14.918810Z","iopub.status.idle":"2025-04-10T12:47:15.209758Z","shell.execute_reply.started":"2025-04-10T12:47:14.918775Z","shell.execute_reply":"2025-04-10T12:47:15.209053Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Aggregating per Audio File:","metadata":{}},{"cell_type":"code","source":"# -----------------------------\n# File-Level Aggregation\n# -----------------------------\ntest_preds_proba = model.predict_proba(test_X)\ntest_features['pred'] = test_preds\ntest_features['true'] = test_y\n\nfile_votes = defaultdict(list)\nfile_true = {}\n\nfor _, row in test_features.iterrows():\n    fname = row['filename']\n    file_votes[fname].append(row['pred'])\n    file_true[fname] = row['true']\n\nfinal_preds, final_labels = [], []\n\nfor fname, preds in file_votes.items():\n    vote = Counter(preds).most_common(1)[0][0]\n    final_preds.append(vote)\n    final_labels.append(file_true[fname])\n\nfile_level_acc = accuracy_score(final_labels, final_preds)\nprint(\"✅ File-level Accuracy (Majority Vote):\", file_level_acc)\nprint(\"✅ File-level Classification Report:\\n\", classification_report(final_labels, final_preds, target_names=label_encoder.classes_))\n\nmlflow.log_metric(\"file_level_accuracy\", file_level_acc)\n\n# File-level confusion matrix\ncm = confusion_matrix(final_labels, final_preds)\nplt.figure(figsize=(6, 5))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\",\n            xticklabels=label_encoder.classes_,\n            yticklabels=label_encoder.classes_)\nplt.title(\"File-level Confusion Matrix (XGBoost)\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.tight_layout()\n\nwith tempfile.NamedTemporaryFile(suffix=\".png\", delete=False) as tmp:\n    plt.savefig(tmp.name)\n    mlflow.log_artifact(tmp.name, artifact_path=\"file_level_confusion_matrix\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-10T12:47:15.210630Z","iopub.execute_input":"2025-04-10T12:47:15.210872Z","iopub.status.idle":"2025-04-10T12:47:15.607752Z","shell.execute_reply.started":"2025-04-10T12:47:15.210851Z","shell.execute_reply":"2025-04-10T12:47:15.606834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Regnets:\n\n### First we will extract spectograms:","metadata":{}},{"cell_type":"code","source":"import os\nimport librosa\nimport numpy as np\nimport cv2\nfrom tqdm import tqdm\n\ndef mmss_to_seconds(mmss: str):\n    mins, secs = map(int, mmss.split(\":\"))\n    return mins * 60 + secs\n\ndef normalize(x):\n    return (x - x.min()) / (x.max() - x.min() + 1e-8)\n\ndef resize_feat(feat, target_shape):\n    from cv2 import resize, INTER_LINEAR\n    return resize(feat, target_shape, interpolation=INTER_LINEAR)\n\ndef extract_rgb_features(y, sr, img_size=(224, 224)):\n    try:\n        # Compute features\n        mel = librosa.feature.melspectrogram(y=y, sr=sr, n_mels=128)\n        log_mel = librosa.power_to_db(mel, ref=np.max)\n        delta = librosa.feature.delta(log_mel)\n        chroma = librosa.feature.chroma_cqt(y=y, sr=sr)\n\n        # Resize features\n        log_mel_resized = resize_feat(log_mel, img_size)\n        delta_resized = resize_feat(delta, img_size)\n        chroma_resized = resize_feat(chroma, img_size)\n\n        # Normalize\n        log_mel_norm = normalize(log_mel_resized)\n        delta_norm = normalize(delta_resized)\n        chroma_norm = normalize(chroma_resized)\n\n        # Stack to form RGB image\n        rgb_image = np.stack([log_mel_norm, delta_norm, chroma_norm], axis=-1)\n        rgb_image_uint8 = (rgb_image * 255).astype(np.uint8)\n\n        return rgb_image_uint8\n\n    except Exception as e:\n        print(f\"[ERROR in feature extraction]: {e}\")\n        return None\n\ndef save_chunked_rgb_images_trimmed(df, speech_info, audio_root, output_dir, chunk_duration=5, img_size=(224, 224)):\n    os.makedirs(output_dir, exist_ok=True)\n    speech_map = {entry['filename']: entry['speech_timestamps'] for entry in speech_info}\n\n    for _, row in tqdm(df.iterrows(), total=len(df)):\n        filename = row['filename']\n        animal_class = row['animal_class']\n        input_path = os.path.join(audio_root, filename)\n\n        try:\n            y, sr = librosa.load(input_path, sr=None)\n        except Exception as e:\n            print(f\"[ERROR] loading {filename}: {e}\")\n            continue\n\n        duration_sec = len(y) / sr\n\n        # Trim human speech if found\n        if filename in speech_map:\n            speech = speech_map[filename]\n            if speech:\n                speech_starts = [mmss_to_seconds(seg['start']) for seg in speech]\n                speech_ends = [mmss_to_seconds(seg['end']) for seg in speech]\n                earliest = min(speech_starts)\n                latest = max(speech_ends)\n\n                if latest < duration_sec / 2:\n                    trim_sec = int(np.ceil(latest / chunk_duration) * chunk_duration)\n                    y = y[int(trim_sec * sr):]\n                elif earliest > duration_sec / 2:\n                    trim_sec = int(np.floor(earliest / chunk_duration) * chunk_duration)\n                    y = y[:int(trim_sec * sr)]\n                else:\n                    # Mid-speech → skip\n                    pass\n\n        # 5-second chunking\n        samples_per_chunk = int(sr * chunk_duration)\n        n_chunks = int(len(y) / samples_per_chunk)\n        if n_chunks == 0:\n            pass\n\n        base_name = os.path.splitext(os.path.basename(filename))[0]\n        class_dir = os.path.join(output_dir, animal_class)\n        os.makedirs(class_dir, exist_ok=True)\n\n        for i in range(n_chunks):\n            start = i * samples_per_chunk\n            end = start + samples_per_chunk\n            chunk = y[start:end]\n\n            rgb_img = extract_rgb_features(chunk, sr, img_size=img_size)\n            if rgb_img is None:\n                continue\n\n            out_name = f\"{base_name}_clip_{i}.png\"\n            out_path = os.path.join(class_dir, out_name)\n\n            try:\n                cv2.imwrite(out_path, rgb_img)\n            except Exception as e:\n                print(f\"[ERROR] saving {out_path}: {e}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:32:10.307961Z","iopub.execute_input":"2025-04-14T13:32:10.308284Z","iopub.status.idle":"2025-04-14T13:32:10.644959Z","shell.execute_reply.started":"2025-04-14T13:32:10.308254Z","shell.execute_reply":"2025-04-14T13:32:10.644270Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"save_chunked_rgb_images_trimmed(train_df, speech_info_train, audio_root='/kaggle/input/birdclef-2025/train_audio', output_dir='/kaggle/working/spectrogram_chunks/train_chunked')\nsave_chunked_rgb_images_trimmed(val_df, speech_info_val, audio_root='/kaggle/input/birdclef-2025/train_audio', output_dir='/kaggle/working/spectrogram_chunks/val_chunked')\nsave_chunked_rgb_images_trimmed(test_df, speech_info_test, audio_root='/kaggle/input/birdclef-2025/train_audio', output_dir='/kaggle/working/spectrogram_chunks/test_chunked')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:32:32.205348Z","iopub.execute_input":"2025-04-14T13:32:32.205762Z","iopub.status.idle":"2025-04-14T13:43:35.454375Z","shell.execute_reply.started":"2025-04-14T13:32:32.205726Z","shell.execute_reply":"2025-04-14T13:43:35.453479Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Dataset & Dataloader:","metadata":{}},{"cell_type":"code","source":"import torchvision.transforms as T\nfrom torchvision.datasets import ImageFolder\nfrom torch.utils.data import DataLoader\n\n# Transforms\ntransform = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(mean=[0.5]*3, std=[0.5]*3)\n])\n\n# Datasets\ntrain_ds = ImageFolder(\"/kaggle/working/spectrogram_chunks/train_chunked\", transform=transform)\nval_ds   = ImageFolder(\"/kaggle/working/spectrogram_chunks/val_chunked\", transform=transform)\ntest_ds  = ImageFolder(\"/kaggle/working/spectrogram_chunks/test_chunked\", transform=transform)\n\n# Dataloaders\ntrain_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=2)\nval_loader   = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=2)\ntest_loader  = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:43:35.455621Z","iopub.execute_input":"2025-04-14T13:43:35.455979Z","iopub.status.idle":"2025-04-14T13:43:35.980345Z","shell.execute_reply.started":"2025-04-14T13:43:35.455940Z","shell.execute_reply":"2025-04-14T13:43:35.979319Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Regnet Model:","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport pytorch_lightning as pl\nimport torchvision.models as models\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import StepLR\nclass RegNetClassifier(pl.LightningModule):\n    def __init__(self, num_classes=4, lr=1e-4):\n        super().__init__()\n        self.save_hyperparameters()\n\n        # RegNet setup\n        self.model = models.regnet_y_400mf(pretrained=True)\n        self.model.fc = nn.Linear(self.model.fc.in_features, num_classes)\n        self.criterion = nn.CrossEntropyLoss()\n\n        # For manual tracking\n        self.val_losses = []\n        self.val_accuracies = []\n        self._val_loss_batches = []\n        self._val_acc_batches = []\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.criterion(logits, y)\n        acc = (logits.argmax(dim=1) == y).float().mean()\n        self.log(\"train_loss\", loss)\n        self.log(\"train_acc\", acc, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.criterion(logits, y)\n        acc = (logits.argmax(dim=1) == y).float().mean()\n\n        # Save to temporary lists for epoch-end aggregation\n        self._val_loss_batches.append(loss)\n        self._val_acc_batches.append(acc)\n\n        # Log per batch for trainer bar\n        self.log(\"val_loss\", loss, prog_bar=True, on_step=False, on_epoch=True)\n        self.log(\"val_acc\", acc, prog_bar=True, on_step=False, on_epoch=True)\n\n    def on_validation_epoch_end(self):\n        # Average batch metrics for this epoch\n        avg_loss = torch.stack(self._val_loss_batches).mean().item()\n        avg_acc = torch.stack(self._val_acc_batches).mean().item()\n\n        # Store for plotting\n        self.val_losses.append(avg_loss)\n        self.val_accuracies.append(avg_acc)\n\n        # Log epoch-level values\n        self.log(\"val_loss\", avg_loss, prog_bar=True)\n        self.log(\"val_acc\", avg_acc, prog_bar=True)\n\n        # Clear lists for next epoch\n        self._val_loss_batches.clear()\n        self._val_acc_batches.clear()\n\n    def configure_optimizers(self):\n        optimizer = optim.Adam(self.parameters(), lr=self.hparams.lr)\n        scheduler = StepLR(optimizer, step_size=2, gamma=0.8)\n        return [optimizer], [scheduler]\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:43:35.981461Z","iopub.execute_input":"2025-04-14T13:43:35.981766Z","iopub.status.idle":"2025-04-14T13:43:35.991685Z","shell.execute_reply.started":"2025-04-14T13:43:35.981739Z","shell.execute_reply":"2025-04-14T13:43:35.990760Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pytorch_lightning.loggers import MLFlowLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nfrom pytorch_lightning import Trainer\n\nmlf_logger = MLFlowLogger(\n    experiment_name=\"BirdCLEF-RegNet\",\n    tracking_uri=\"file:/kaggle/working/mlruns\"  # or another location if remote\n)\n\ncheckpoint = ModelCheckpoint(monitor=\"val_acc\", mode=\"max\", save_top_k=1)\n\nmodel = RegNetClassifier(num_classes=4, lr=1e-3)\n\ntrainer = Trainer(\n    max_epochs=10,\n    accelerator=\"auto\",\n    callbacks=[checkpoint],\n    logger=mlf_logger\n)\n\ntrainer.fit(model, train_loader, val_loader)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:43:35.992884Z","iopub.execute_input":"2025-04-14T13:43:35.993221Z","iopub.status.idle":"2025-04-14T13:45:58.574557Z","shell.execute_reply.started":"2025-04-14T13:43:35.993193Z","shell.execute_reply":"2025-04-14T13:45:58.573721Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Convert to arrays for smooth plotting\nval_losses = model.val_losses\nval_accuracies = model.val_accuracies\nval_losses = model.val_losses[1:]\nval_accuracies = model.val_accuracies[1:]\n\nepochs = range(1, len(val_losses) + 1)\n\nplt.figure(figsize=(12, 5))\n\n# Plot Validation Loss\nplt.subplot(1, 2, 1)\nplt.plot(epochs, val_losses, 'o-', label='Validation Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.title('Validation Loss over Epochs')\nplt.grid(True)\nplt.legend()\n\n# Plot Validation Accuracy\nplt.subplot(1, 2, 2)\nplt.plot(epochs, val_accuracies, 'o-', label='Validation Accuracy')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.title('Validation Accuracy over Epochs')\nplt.grid(True)\nplt.legend()\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:45:58.575622Z","iopub.execute_input":"2025-04-14T13:45:58.575890Z","iopub.status.idle":"2025-04-14T13:45:58.987698Z","shell.execute_reply.started":"2025-04-14T13:45:58.575866Z","shell.execute_reply":"2025-04-14T13:45:58.986881Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mlflow.end_run()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:45:58.989606Z","iopub.execute_input":"2025-04-14T13:45:58.989894Z","iopub.status.idle":"2025-04-14T13:45:58.993751Z","shell.execute_reply.started":"2025-04-14T13:45:58.989871Z","shell.execute_reply":"2025-04-14T13:45:58.992743Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Test the Model:","metadata":{}},{"cell_type":"code","source":"import mlflow\nimport tempfile\nfrom sklearn.metrics import accuracy_score, classification_report, confusion_matrix\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom collections import defaultdict\nimport numpy as np\nfrom scipy.stats import mode\n\nwith mlflow.start_run(run_name=\"RegNet_Eval\"):\n    model.eval()\n    all_preds, all_labels = [], []\n\n    with torch.no_grad():\n        for x, y in test_loader:\n            x = x.to(model.device)\n            logits = model(x)\n            preds = logits.argmax(dim=1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(y.numpy())\n\n    acc = accuracy_score(all_labels, all_preds)\n    report = classification_report(all_labels, all_preds, target_names=test_ds.classes)\n\n    print(\"✅ Test Accuracy:\", acc)\n    print(report)\n\n    mlflow.log_metric(\"test_accuracy\", acc)\n\n    # Log confusion matrix image\n    cm = confusion_matrix(all_labels, all_preds)\n    plt.figure(figsize=(6, 5))\n    sns.heatmap(cm, annot=True, fmt=\"d\", xticklabels=test_ds.classes, yticklabels=test_ds.classes, cmap=\"Greens\")\n    plt.title(\"Test Confusion Matrix\")\n    plt.xlabel(\"Predicted\")\n    plt.ylabel(\"True\")\n    plt.tight_layout()\n\n    with tempfile.NamedTemporaryFile(suffix=\".png\", delete=False) as tmp:\n        plt.savefig(tmp.name)\n        mlflow.log_artifact(tmp.name, artifact_path=\"test_confusion_matrix\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:45:58.994891Z","iopub.execute_input":"2025-04-14T13:45:58.995146Z","iopub.status.idle":"2025-04-14T13:46:15.239945Z","shell.execute_reply.started":"2025-04-14T13:45:58.995126Z","shell.execute_reply":"2025-04-14T13:46:15.238922Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Aggregating per audio file:","metadata":{}},{"cell_type":"code","source":"file_probs_model = defaultdict(list)\nfile_targets = {}\nall_image_paths = [s[0] for s in test_loader.dataset.samples]\nsample_idx = 0\n\nwith torch.no_grad():\n    for batch, labels in test_loader:\n        batch = batch.to(model.device)\n        logits = model(batch)\n        probs = torch.softmax(logits, dim=1).cpu().numpy()\n\n        for j in range(len(batch)):\n            img_path = all_image_paths[sample_idx]\n            true_class = test_loader.dataset.samples[sample_idx][1]\n            base_filename = os.path.basename(img_path).split(\"_clip\")[0]\n\n            file_probs_model[base_filename].append(probs[j])\n            file_targets[base_filename] = true_class\n            sample_idx += 1\n\nfinal_preds, final_labels = [], []\n\nfor fname, prob_list in file_probs_model.items():\n    votes = [np.argmax(p) for p in prob_list]\n    pred = mode(votes, keepdims=True).mode[0]\n    final_preds.append(pred)\n    final_labels.append(file_targets[fname])\n\nfile_acc = accuracy_score(final_labels, final_preds)\nprint(\"✅ File-level Accuracy:\", file_acc)\nprint(\"✅ File-level Classification Report:\")\nprint(classification_report(final_labels, final_preds, target_names=test_ds.classes))\n\nmlflow.log_metric(\"file_level_accuracy\", file_acc)\n\n# File-level confusion matrix\ncm = confusion_matrix(final_labels, final_preds)\nplt.figure(figsize=(6, 5))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n            xticklabels=test_ds.classes,\n            yticklabels=test_ds.classes)\nplt.title(\"File-level Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.tight_layout()\n\nwith tempfile.NamedTemporaryFile(suffix=\".png\", delete=False) as tmp:\n    plt.savefig(tmp.name)\n    mlflow.log_artifact(tmp.name, artifact_path=\"file_level_confusion_matrix\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:46:15.240974Z","iopub.execute_input":"2025-04-14T13:46:15.241367Z","iopub.status.idle":"2025-04-14T13:46:31.195314Z","shell.execute_reply.started":"2025-04-14T13:46:15.241330Z","shell.execute_reply":"2025-04-14T13:46:31.194545Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mlflow.end_run()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T14:01:11.938934Z","iopub.execute_input":"2025-04-14T14:01:11.939262Z","iopub.status.idle":"2025-04-14T14:01:11.945460Z","shell.execute_reply.started":"2025-04-14T14:01:11.939235Z","shell.execute_reply":"2025-04-14T14:01:11.944781Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Efficientnet Model:","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport pytorch_lightning as pl\nimport torchvision.models as models\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import StepLR\n\nclass EfficientNetClassifier(pl.LightningModule):\n    def __init__(self, num_classes=4, lr=1e-4):\n        super().__init__()\n        self.save_hyperparameters()\n\n        # Load pretrained EfficientNet \n        self.model = models.efficientnet_b3(pretrained=True)\n        self.model.classifier[1] = nn.Linear(self.model.classifier[1].in_features, num_classes)\n\n        self.criterion = nn.CrossEntropyLoss()\n\n        # Manual metric tracking\n        self.val_losses = []\n        self.val_accuracies = []\n        self._val_loss_batches = []\n        self._val_acc_batches = []\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.criterion(logits, y)\n        acc = (logits.argmax(dim=1) == y).float().mean()\n        self.log(\"train_loss\", loss)\n        self.log(\"train_acc\", acc, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.criterion(logits, y)\n        acc = (logits.argmax(dim=1) == y).float().mean()\n\n        # Store per batch\n        self._val_loss_batches.append(loss)\n        self._val_acc_batches.append(acc)\n\n        self.log(\"val_loss\", loss, prog_bar=True, on_step=False, on_epoch=True)\n        self.log(\"val_acc\", acc, prog_bar=True, on_step=False, on_epoch=True)\n\n    def on_validation_epoch_end(self):\n        avg_loss = torch.stack(self._val_loss_batches).mean().item()\n        avg_acc = torch.stack(self._val_acc_batches).mean().item()\n\n        self.val_losses.append(avg_loss)\n        self.val_accuracies.append(avg_acc)\n\n        self.log(\"val_loss\", avg_loss, prog_bar=True)\n        self.log(\"val_acc\", avg_acc, prog_bar=True)\n\n        self._val_loss_batches.clear()\n        self._val_acc_batches.clear()\n\n    def configure_optimizers(self):\n        optimizer = optim.Adam(self.parameters(), lr=self.hparams.lr)\n        scheduler = StepLR(optimizer, step_size=2, gamma=0.8)\n        return [optimizer], [scheduler]\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:46:48.269088Z","iopub.execute_input":"2025-04-14T13:46:48.269381Z","iopub.status.idle":"2025-04-14T13:46:48.279234Z","shell.execute_reply.started":"2025-04-14T13:46:48.269357Z","shell.execute_reply":"2025-04-14T13:46:48.278286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pytorch_lightning.loggers import MLFlowLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nfrom pytorch_lightning import Trainer\n\nmlf_logger = MLFlowLogger(\n    experiment_name=\"BirdCLEF-EfficientNet\",\n    tracking_uri=\"file:/kaggle/working/mlruns\"  # or another location if remote\n)\n\ncheckpoint = ModelCheckpoint(monitor=\"val_acc\", mode=\"max\", save_top_k=1)\n\neff_net_model = EfficientNetClassifier(num_classes=4, lr=1e-3)\n\ntrainer = Trainer(\n    max_epochs=10,\n    accelerator=\"auto\",\n    callbacks=[checkpoint],\n    logger=mlf_logger\n)\n\ntrainer.fit(eff_net_model, train_loader, val_loader)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:46:52.995498Z","iopub.execute_input":"2025-04-14T13:46:52.995835Z","iopub.status.idle":"2025-04-14T13:53:01.544218Z","shell.execute_reply.started":"2025-04-14T13:46:52.995777Z","shell.execute_reply":"2025-04-14T13:53:01.543424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Convert to arrays for smooth plotting\nval_losses = eff_net_model.val_losses\nval_accuracies = eff_net_model.val_accuracies\nval_losses = eff_net_model.val_losses[1:]\nval_accuracies = eff_net_model.val_accuracies[1:]\n\nepochs = range(1, len(val_losses) + 1)\n\nplt.figure(figsize=(12, 5))\n\n# Plot Validation Loss\nplt.subplot(1, 2, 1)\nplt.plot(epochs, val_losses, 'o-', label='Validation Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.title('Validation Loss over Epochs')\nplt.grid(True)\nplt.legend()\n\n# Plot Validation Accuracy\nplt.subplot(1, 2, 2)\nplt.plot(epochs, val_accuracies, 'o-', label='Validation Accuracy')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.title('Validation Accuracy over Epochs')\nplt.grid(True)\nplt.legend()\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:53:01.545723Z","iopub.execute_input":"2025-04-14T13:53:01.546114Z","iopub.status.idle":"2025-04-14T13:53:01.934440Z","shell.execute_reply.started":"2025-04-14T13:53:01.546074Z","shell.execute_reply":"2025-04-14T13:53:01.933607Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Test the Model:","metadata":{}},{"cell_type":"code","source":"import mlflow\nimport tempfile\nfrom sklearn.metrics import accuracy_score, classification_report, confusion_matrix\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom collections import defaultdict\nimport numpy as np\nfrom scipy.stats import mode\n\nwith mlflow.start_run(run_name=\"EfficientNet_Eval\"):\n    eff_net_model.eval()\n    all_preds, all_labels = [], []\n\n    with torch.no_grad():\n        for x, y in test_loader:\n            x = x.to(eff_net_model.device)\n            logits = eff_net_model(x)\n            preds = logits.argmax(dim=1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(y.numpy())\n\n    acc = accuracy_score(all_labels, all_preds)\n    report = classification_report(all_labels, all_preds, target_names=test_ds.classes)\n\n    print(\"✅ Test Accuracy:\", acc)\n    print(report)\n\n    mlflow.log_metric(\"test_accuracy\", acc)\n\n    # Log confusion matrix image\n    cm = confusion_matrix(all_labels, all_preds)\n    plt.figure(figsize=(6, 5))\n    sns.heatmap(cm, annot=True, fmt=\"d\", xticklabels=test_ds.classes, yticklabels=test_ds.classes, cmap=\"Greens\")\n    plt.title(\"Test Confusion Matrix\")\n    plt.xlabel(\"Predicted\")\n    plt.ylabel(\"True\")\n    plt.tight_layout()\n\n    with tempfile.NamedTemporaryFile(suffix=\".png\", delete=False) as tmp:\n        plt.savefig(tmp.name)\n        mlflow.log_artifact(tmp.name, artifact_path=\"test_confusion_matrix\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:53:01.936023Z","iopub.execute_input":"2025-04-14T13:53:01.936279Z","iopub.status.idle":"2025-04-14T13:54:13.790025Z","shell.execute_reply.started":"2025-04-14T13:53:01.936256Z","shell.execute_reply":"2025-04-14T13:54:13.789166Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Aggregating per audio file:","metadata":{}},{"cell_type":"code","source":"file_probs_eff = defaultdict(list)\nfile_targets = {}\nall_image_paths = [s[0] for s in test_loader.dataset.samples]\nsample_idx = 0\n\nwith torch.no_grad():\n    for batch, labels in test_loader:\n        batch = batch.to(eff_net_model.device)\n        logits = eff_net_model(batch)\n        probs = torch.softmax(logits, dim=1).cpu().numpy()\n\n        for j in range(len(batch)):\n            img_path = all_image_paths[sample_idx]\n            true_class = test_loader.dataset.samples[sample_idx][1]\n            base_filename = os.path.basename(img_path).split(\"_clip\")[0]\n\n            file_probs_eff[base_filename].append(probs[j])\n            file_targets[base_filename] = true_class\n            sample_idx += 1\n\nfinal_preds, final_labels = [], []\n\nfor fname, prob_list in file_probs_eff.items():\n    votes = [np.argmax(p) for p in prob_list]\n    pred = mode(votes, keepdims=True).mode[0]\n    final_preds.append(pred)\n    final_labels.append(file_targets[fname])\n\nfile_acc = accuracy_score(final_labels, final_preds)\nprint(\"✅ File-level Accuracy:\", file_acc)\nprint(\"✅ File-level Classification Report:\")\nprint(classification_report(final_labels, final_preds, target_names=test_ds.classes))\n\nmlflow.log_metric(\"file_level_accuracy\", file_acc)\n\n# File-level confusion matrix\ncm = confusion_matrix(final_labels, final_preds)\nplt.figure(figsize=(6, 5))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n            xticklabels=test_ds.classes,\n            yticklabels=test_ds.classes)\nplt.title(\"File-level Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.tight_layout()\n\nwith tempfile.NamedTemporaryFile(suffix=\".png\", delete=False) as tmp:\n    plt.savefig(tmp.name)\n    mlflow.log_artifact(tmp.name, artifact_path=\"file_level_confusion_matrix\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:54:13.791047Z","iopub.execute_input":"2025-04-14T13:54:13.791309Z","iopub.status.idle":"2025-04-14T13:55:24.165044Z","shell.execute_reply.started":"2025-04-14T13:54:13.791284Z","shell.execute_reply":"2025-04-14T13:55:24.164015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mlflow.end_run()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T14:14:13.117172Z","iopub.execute_input":"2025-04-14T14:14:13.117495Z","iopub.status.idle":"2025-04-14T14:14:13.123576Z","shell.execute_reply.started":"2025-04-14T14:14:13.117472Z","shell.execute_reply":"2025-04-14T14:14:13.122619Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Extra Context(2.5sec before, 2.5sec after):","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nfrom torch.utils.data import Dataset\nfrom torchvision import transforms\nfrom PIL import Image\nimport numpy as np\n\nclass SpectrogramContextDataset(Dataset):\n    def __init__(self, image_dir, transform=None):\n        self.image_paths = []\n        self.labels = []\n        self.filenames = []\n\n        for class_name in sorted(os.listdir(image_dir)):\n            class_dir = os.path.join(image_dir, class_name)\n            if not os.path.isdir(class_dir):\n                continue\n            for file in sorted(os.listdir(class_dir)):  # ensures chronological chunk order\n                if file.endswith('.png'):\n                    self.image_paths.append(os.path.join(class_dir, file))\n                    self.labels.append(class_name)\n                    self.filenames.append(file.split(\"_clip\")[0])  # e.g., iNat1122209\n\n        self.classes = sorted(set(self.labels))\n        self.class_to_idx = {c: i for i, c in enumerate(self.classes)}\n        self.transform = transform\n        self.samples = list(zip(self.image_paths, self.labels))\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def _load_half(self, path, side='left'):\n        img = Image.open(path).convert(\"RGB\")\n        img = np.array(img)\n        h, w, c = img.shape\n        if side == 'left':\n            half = img[:, :w // 2, :]\n        else:\n            half = img[:, w // 2:, :]\n        return half  # shape: (224, 112, 3)\n\n    def __getitem__(self, idx):\n        center_path = self.image_paths[idx]\n        center_img = Image.open(center_path).convert(\"RGB\")\n        center_arr = np.array(center_img)  # shape: (224, 224, 3)\n        fname = self.filenames[idx]\n\n        # Default paddings\n        left_half = np.zeros((224, 112, 3), dtype=np.uint8)\n        right_half = np.zeros((224, 112, 3), dtype=np.uint8)\n\n        # LEFT: prev image from same file\n        if idx > 0 and self.filenames[idx - 1] == fname:\n            left_half = self._load_half(self.image_paths[idx - 1], side='right')\n\n        # RIGHT: next image from same file\n        if idx < len(self.image_paths) - 1 and self.filenames[idx + 1] == fname:\n            right_half = self._load_half(self.image_paths[idx + 1], side='left')\n\n        # Combine horizontally: (224, 448, 3)\n        full_img = np.concatenate([left_half, center_arr, right_half], axis=1)\n        full_img = Image.fromarray(full_img.astype(np.uint8))\n\n        if self.transform:\n            full_img = self.transform(full_img)\n\n        label = self.class_to_idx[self.labels[idx]]\n        return full_img, label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:55:45.398603Z","iopub.execute_input":"2025-04-14T13:55:45.399115Z","iopub.status.idle":"2025-04-14T13:55:45.410871Z","shell.execute_reply.started":"2025-04-14T13:55:45.399075Z","shell.execute_reply":"2025-04-14T13:55:45.409893Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision import transforms\nfrom torch.utils.data import DataLoader\n\n# Path to your context-aware spectrogram folders\ntrain_dir = \"/kaggle/working/spectrogram_chunks/train_chunked\"\nval_dir   = \"/kaggle/working/spectrogram_chunks/val_chunked\"\ntest_dir  = \"/kaggle/working/spectrogram_chunks/test_chunked\"\n\n# Shared transform\ntransform = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.5]*3, std=[0.5]*3)\n])\n\n# Context-aware datasets\ntrain_ds = SpectrogramContextDataset(train_dir, transform=transform)\nval_ds   = SpectrogramContextDataset(val_dir, transform=transform)\ntest_ds  = SpectrogramContextDataset(test_dir, transform=transform)\n\n# Dataloaders\ntrain_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=2)\nval_loader   = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=2)\ntest_loader  = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:55:48.887888Z","iopub.execute_input":"2025-04-14T13:55:48.888192Z","iopub.status.idle":"2025-04-14T13:55:48.908749Z","shell.execute_reply.started":"2025-04-14T13:55:48.888169Z","shell.execute_reply":"2025-04-14T13:55:48.907851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport pytorch_lightning as pl\nimport torchvision.models as models\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import StepLR\n\nclass RegNetClassifier(pl.LightningModule):\n    def __init__(self, num_classes=4, lr=1e-4):\n        super().__init__()\n        self.save_hyperparameters()\n\n        # Load RegNet with wide support\n        self.model = models.regnet_y_400mf(pretrained=True)\n        self.model.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n        self.model.fc = nn.Linear(self.model.fc.in_features, num_classes)\n\n        self.criterion = nn.CrossEntropyLoss()\n\n        # Manual metric tracking\n        self.val_losses = []\n        self.val_accuracies = []\n        self._val_loss_batches = []\n        self._val_acc_batches = []\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.criterion(logits, y)\n        acc = (logits.argmax(dim=1) == y).float().mean()\n        self.log(\"train_loss\", loss)\n        self.log(\"train_acc\", acc, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.criterion(logits, y)\n        acc = (logits.argmax(dim=1) == y).float().mean()\n\n        self._val_loss_batches.append(loss)\n        self._val_acc_batches.append(acc)\n\n        self.log(\"val_loss\", loss, prog_bar=True, on_step=False, on_epoch=True)\n        self.log(\"val_acc\", acc, prog_bar=True, on_step=False, on_epoch=True)\n\n    def on_validation_epoch_end(self):\n        avg_loss = torch.stack(self._val_loss_batches).mean().item()\n        avg_acc = torch.stack(self._val_acc_batches).mean().item()\n\n        self.val_losses.append(avg_loss)\n        self.val_accuracies.append(avg_acc)\n\n        self.log(\"val_loss\", avg_loss, prog_bar=True)\n        self.log(\"val_acc\", avg_acc, prog_bar=True)\n\n        self._val_loss_batches.clear()\n        self._val_acc_batches.clear()\n\n    def configure_optimizers(self):\n        optimizer = optim.Adam(self.parameters(), lr=self.hparams.lr)\n        scheduler = StepLR(optimizer, step_size=2, gamma=0.8)\n        return [optimizer], [scheduler]\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:55:52.309748Z","iopub.execute_input":"2025-04-14T13:55:52.310105Z","iopub.status.idle":"2025-04-14T13:55:52.320401Z","shell.execute_reply.started":"2025-04-14T13:55:52.310077Z","shell.execute_reply":"2025-04-14T13:55:52.319511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pytorch_lightning import Trainer\nfrom pytorch_lightning.callbacks import ModelCheckpoint\n\ncheckpoint_cb = ModelCheckpoint(monitor=\"val_acc\", mode=\"max\", save_top_k=1)\n\nmodel_regnet_extra = RegNetClassifier(num_classes=4)\n\ntrainer = Trainer(\n    max_epochs=10,\n    accelerator=\"auto\",\n    callbacks=[checkpoint_cb]\n)\n\ntrainer.fit(model_regnet_extra, train_loader, val_loader)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T13:55:55.546579Z","iopub.execute_input":"2025-04-14T13:55:55.547006Z","iopub.status.idle":"2025-04-14T14:00:50.591610Z","shell.execute_reply.started":"2025-04-14T13:55:55.546968Z","shell.execute_reply":"2025-04-14T14:00:50.590784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Convert to arrays for smooth plotting\nval_losses = model_regnet_extra.val_losses\nval_accuracies = model_regnet_extra.val_accuracies\nval_losses = model_regnet_extra.val_losses[1:]\nval_accuracies = model_regnet_extra.val_accuracies[1:]\n\nepochs = range(1, len(val_losses) + 1)\n\nplt.figure(figsize=(12, 5))\n\n# Plot Validation Loss\nplt.subplot(1, 2, 1)\nplt.plot(epochs, val_losses, 'o-', label='Validation Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.title('Validation Loss over Epochs')\nplt.grid(True)\nplt.legend()\n\n# Plot Validation Accuracy\nplt.subplot(1, 2, 2)\nplt.plot(epochs, val_accuracies, 'o-', label='Validation Accuracy')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.title('Validation Accuracy over Epochs')\nplt.grid(True)\nplt.legend()\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T14:00:50.593098Z","iopub.execute_input":"2025-04-14T14:00:50.593403Z","iopub.status.idle":"2025-04-14T14:00:50.996598Z","shell.execute_reply.started":"2025-04-14T14:00:50.593375Z","shell.execute_reply":"2025-04-14T14:00:50.995734Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import mlflow\nimport tempfile\nfrom sklearn.metrics import accuracy_score, classification_report, confusion_matrix\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom collections import defaultdict\nimport numpy as np\nfrom scipy.stats import mode\n\nwith mlflow.start_run(run_name=\"RegNet_Eval_extra_context\"):\n    model_regnet_extra.eval()\n    all_preds, all_labels = [], []\n\n    with torch.no_grad():\n        for x, y in test_loader:\n            x = x.to(model_regnet_extra.device)\n            logits = model_regnet_extra(x)\n            preds = logits.argmax(dim=1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(y.numpy())\n\n    acc = accuracy_score(all_labels, all_preds)\n    report = classification_report(all_labels, all_preds, target_names=test_ds.classes)\n\n    print(\"✅ Test Accuracy:\", acc)\n    print(report)\n\n    mlflow.log_metric(\"test_accuracy\", acc)\n\n    # Log confusion matrix image\n    cm = confusion_matrix(all_labels, all_preds)\n    plt.figure(figsize=(6, 5))\n    sns.heatmap(cm, annot=True, fmt=\"d\", xticklabels=test_ds.classes, yticklabels=test_ds.classes, cmap=\"Greens\")\n    plt.title(\"Test Confusion Matrix\")\n    plt.xlabel(\"Predicted\")\n    plt.ylabel(\"True\")\n    plt.tight_layout()\n\n    with tempfile.NamedTemporaryFile(suffix=\".png\", delete=False) as tmp:\n        plt.savefig(tmp.name)\n        mlflow.log_artifact(tmp.name, artifact_path=\"test_confusion_matrix\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T14:14:20.527097Z","iopub.execute_input":"2025-04-14T14:14:20.527404Z","iopub.status.idle":"2025-04-14T14:14:58.761698Z","shell.execute_reply.started":"2025-04-14T14:14:20.527382Z","shell.execute_reply":"2025-04-14T14:14:58.760831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Map of RegNet predictions per file\nfile_probs_extra = defaultdict(list)\nfile_targets = {}\n\n# Get filepaths in order\nall_image_paths = [s[0] for s in test_loader.dataset.samples]\nsample_idx = 0\n\n# Mapping from class index to label\nidx_to_class = {i: c for i, c in enumerate(test_ds.classes)}\n\nwith torch.no_grad():\n    for batch, labels in test_loader:\n        batch = batch.to(model_regnet_extra.device)\n        logits = model_regnet_extra(batch)\n        probs = torch.softmax(logits, dim=1).cpu().numpy()\n\n        for j in range(len(batch)):\n            img_path = all_image_paths[sample_idx]\n            class_idx = test_loader.dataset.samples[sample_idx][1]  # already integer\n            fname = os.path.basename(img_path).split(\"_clip\")[0]\n\n            file_probs_extra[fname].append(probs[j])\n            file_targets[fname] = class_idx  # always integer\n            sample_idx += 1\n\n# Aggregate per file\nfinal_preds, final_labels = [], []\n\nfor fname, prob_list in file_probs_extra.items():\n    vote_indices = [np.argmax(p) for p in prob_list]\n    majority_vote = mode(vote_indices, keepdims=True).mode[0]\n    final_preds.append(majority_vote)\n    final_labels.append(file_targets[fname])\n\n# Accuracy\nfile_acc = accuracy_score(final_labels, final_preds)\nprint(\"✅ File-level Accuracy:\", file_acc)\n\n# Classification Report\nif isinstance(final_labels[0], str):\n    class_to_idx = test_ds.class_to_idx\n    final_labels = [class_to_idx[l] for l in final_labels]\n\nif isinstance(final_preds[0], str):\n    class_to_idx = test_ds.class_to_idx\n    final_preds = [class_to_idx[p] for p in final_preds]\n\nprint(\"✅ File-level Classification Report:\")\nprint(classification_report(final_labels, final_preds, target_names=test_ds.classes))\n\nmlflow.log_metric(\"file_level_accuracy\", file_acc)\n\n# Confusion Matrix\ncm = confusion_matrix(final_labels, final_preds)\nplt.figure(figsize=(6, 5))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n            xticklabels=test_ds.classes,\n            yticklabels=test_ds.classes)\nplt.title(\"File-level Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.tight_layout()\n\n# Log CM to MLflow\nwith tempfile.NamedTemporaryFile(suffix=\".png\", delete=False) as tmp:\n    plt.savefig(tmp.name)\n    mlflow.log_artifact(tmp.name, artifact_path=\"file_level_confusion_matrix\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T14:15:26.360267Z","iopub.execute_input":"2025-04-14T14:15:26.360619Z","iopub.status.idle":"2025-04-14T14:16:04.575271Z","shell.execute_reply.started":"2025-04-14T14:15:26.360587Z","shell.execute_reply":"2025-04-14T14:16:04.574349Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### predictions aggregation:","metadata":{}},{"cell_type":"code","source":"from collections import defaultdict, Counter\nimport numpy as np\nfrom scipy.stats import mode\n\n# Final maps from previous models:\n# file_probs_model, file_probs_eff, file_probs_extra\n# file_targets: {filename: int}\n\n# Voting function across models\nfinal_preds, final_labels = [], []\n\nfor fname in file_targets.keys():\n    # Gather predictions from all 3 models\n    preds_model = [np.argmax(p) for p in file_probs_model[fname]]\n    preds_eff   = [np.argmax(p) for p in file_probs_eff[fname]]\n    preds_extra = [np.argmax(p) for p in file_probs_extra[fname]]\n\n    # Majority vote per chunk list\n    vote = Counter(preds_model + preds_eff + preds_extra).most_common(1)[0][0]\n    \n    final_preds.append(vote)\n    final_labels.append(file_targets[fname])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T14:16:45.907986Z","iopub.execute_input":"2025-04-14T14:16:45.908331Z","iopub.status.idle":"2025-04-14T14:16:45.917265Z","shell.execute_reply.started":"2025-04-14T14:16:45.908305Z","shell.execute_reply":"2025-04-14T14:16:45.916503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.preprocessing import LabelEncoder\nfrom sklearn.metrics import classification_report, confusion_matrix, accuracy_score\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n# Create and fit label encoder using true class names\nlabel_encoder = LabelEncoder()\nlabel_encoder.fit(final_labels)  # uses ['Aves', 'Mammalia', ...]\n\n# Convert int predictions → string predictions using inverse_transform\nfinal_preds_str = label_encoder.inverse_transform(final_preds)\n\n# Now both predictions and labels are strings\nacc = accuracy_score(final_labels, final_preds_str)\nprint(\"✅ Ensemble File-Level Accuracy:\", acc)\n\n# Classification report\nprint(classification_report(final_labels, final_preds_str))\n\n# Confusion matrix (normalized optional)\ncm = confusion_matrix(final_labels, final_preds_str, labels=label_encoder.classes_)\nplt.figure(figsize=(6, 5))\nsns.heatmap(cm, annot=True, fmt='d',\n            xticklabels=label_encoder.classes_,\n            yticklabels=label_encoder.classes_,\n            cmap='Blues')\nplt.title(\"Ensemble Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T14:16:48.392231Z","iopub.execute_input":"2025-04-14T14:16:48.392522Z","iopub.status.idle":"2025-04-14T14:16:48.629106Z","shell.execute_reply.started":"2025-04-14T14:16:48.392500Z","shell.execute_reply":"2025-04-14T14:16:48.628219Z"}},"outputs":[],"execution_count":null}]}