{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"BirdClef 2025 challenge using Pytorch, transfer learning, and hybrid neural networks","metadata":{}},{"cell_type":"markdown","source":"Import Libraries 📚","metadata":{}},{"cell_type":"code","source":"# !pip install --upgrade timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.557848Z","iopub.execute_input":"2025-09-12T19:23:20.558206Z","iopub.status.idle":"2025-09-12T19:23:20.562414Z","shell.execute_reply.started":"2025-09-12T19:23:20.558180Z","shell.execute_reply":"2025-09-12T19:23:20.561496Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport wandb\nimport pandas as pd\nfrom tqdm import tqdm\nimport torch.nn.functional as F\nfrom torch.nn import init\nfrom torch.utils.data import random_split\nfrom torch.utils.data import DataLoader, Dataset, random_split\nimport torchaudio\nimport math, random\nimport torch\nimport torchaudio\nimport torchaudio.transforms as T\nfrom IPython.display import Audio\nimport numpy as np\nimport librosa\nimport librosa.display\nimport IPython.display\nimport matplotlib.pyplot as plt\nimport plotly.graph_objects as go\nfrom PIL import Image\nimport IPython.display as ipd\n\nimport matplotlib as mpl\nimport matplotlib.pylab as ply\nimport ipywidgets as widgets\nimport seaborn as sns\n\n\nfrom itertools import cycle\n# Set interactive backend\n%matplotlib inline\n\n\ncmap = mpl.cm.get_cmap('coolwarm')\nsns.set_theme(style=\"white\", palette=None)\ncolor_pal = ply.rcParams[\"axes.prop_cycle\"].by_key()[\"color\"]\ncolor_cycle = cycle(ply.rcParams[\"axes.prop_cycle\"].by_key()[\"color\"])","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.564160Z","iopub.execute_input":"2025-09-12T19:23:20.564471Z","iopub.status.idle":"2025-09-12T19:23:20.586575Z","shell.execute_reply.started":"2025-09-12T19:23:20.564444Z","shell.execute_reply":"2025-09-12T19:23:20.585552Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Data Exploration 💥💥","metadata":{}},{"cell_type":"code","source":"DATASET_PATH = '/kaggle/input/birdclef-2025'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.588063Z","iopub.execute_input":"2025-09-12T19:23:20.588377Z","iopub.status.idle":"2025-09-12T19:23:20.607446Z","shell.execute_reply.started":"2025-09-12T19:23:20.588351Z","shell.execute_reply":"2025-09-12T19:23:20.606499Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## To handle our settings  and configurations, let's create a class\nclass Config:\n    seed = 42\n    # Input image size and batch size\n    img_size = [128, 384]\n    \n    # Audio duration, sample rate, and length\n    duration = 5 # second\n    sample_rate = 32000\n    audio_len = duration*sample_rate\n    \n    # STFT parameters\n    nfft = 2028\n    window = 2048\n    hop_length = audio_len // (img_size[1] - 1)\n    fmin = 20\n    fmax = 16000\n    \n    #model name\n    preset = 'efficientnetv2_b2_imagenet'\n    class_names = sorted(os.listdir(f'{DATASET_PATH}/train_audio/'))\n    num_classes = len(class_names)\n    class_labels = list(range(num_classes))\n    label2name = dict(zip(class_labels, class_names))\n    name2label = {v:k for k,v in label2name.items()}\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.608480Z","iopub.execute_input":"2025-09-12T19:23:20.608782Z","iopub.status.idle":"2025-09-12T19:23:20.626176Z","shell.execute_reply.started":"2025-09-12T19:23:20.608758Z","shell.execute_reply":"2025-09-12T19:23:20.625303Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Set Seed for Reproducibility","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(Config.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.627867Z","iopub.execute_input":"2025-09-12T19:23:20.628143Z","iopub.status.idle":"2025-09-12T19:23:20.641630Z","shell.execute_reply.started":"2025-09-12T19:23:20.628121Z","shell.execute_reply":"2025-09-12T19:23:20.640739Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Print out the first 5 items in the label2name and name2label dictionaries\nprint(f\"Number of classes: {Config.num_classes}\")\nprint({k: Config.label2name[k] for k in list(Config.label2name)[:5]})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.642611Z","iopub.execute_input":"2025-09-12T19:23:20.642920Z","iopub.status.idle":"2025-09-12T19:23:20.659271Z","shell.execute_reply.started":"2025-09-12T19:23:20.642901Z","shell.execute_reply":"2025-09-12T19:23:20.658278Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Load the dataframe 🔃","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(f'{DATASET_PATH}/train.csv')\ndf['filepath'] = DATASET_PATH + '/train_audio/' + df.filename\ndf['target'] = df.primary_label.map(Config.name2label)\ndf['filename'] = df.filepath.map(lambda x: x.split('/')[-1])\ndf['xc_id'] = df.filepath.map(lambda x: x.split('/')[-1].split('.')[0])\n\n## display a few rows of the dataframe from columns ['scientific_name', 'scientific_name',  'filepath']\ndf = df.sample(frac=1, random_state=Config.seed)\ndf.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.660139Z","iopub.execute_input":"2025-09-12T19:23:20.660451Z","iopub.status.idle":"2025-09-12T19:23:20.892343Z","shell.execute_reply.started":"2025-09-12T19:23:20.660431Z","shell.execute_reply":"2025-09-12T19:23:20.891455Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Explore the largest and smallest class","metadata":{}},{"cell_type":"code","source":"## Display the number of samples per class and save the result in a dictionary\nclass_counts = df.primary_label.value_counts()\nclass_counts = class_counts.sort_index()\nclass_counts\n# ## Save to a csv file\nclass_counts_csv = pd.DataFrame(class_counts.items(), columns=['class', 'count'])\n## Show the largest and smallest classes with the corresponding counts\n# Find the minimum and maximum counts\nmin_count = class_counts.min()\nmax_count = class_counts.max()\n \nprint(f\"Smallest class: {class_counts_csv['class'][class_counts_csv['count'].idxmin()]} {min_count}\")\nprint(f\"Largest class: {class_counts_csv['class'][class_counts_csv['count'].idxmax()]} {max_count}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.893354Z","iopub.execute_input":"2025-09-12T19:23:20.893675Z","iopub.status.idle":"2025-09-12T19:23:20.905023Z","shell.execute_reply.started":"2025-09-12T19:23:20.893632Z","shell.execute_reply":"2025-09-12T19:23:20.903990Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Function to retreive an audio file 🎵\nlibrosa is a python package for music and audio analysis. It provides the building blocks necessary to create music information retrieval systems","metadata":{}},{"cell_type":"code","source":"# Store the sampling rate as `sr`\ndef load_audio(filepath):\n    audio, sr = librosa.load(filepath)\n    return audio, sr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.905904Z","iopub.execute_input":"2025-09-12T19:23:20.906137Z","iopub.status.idle":"2025-09-12T19:23:20.922146Z","shell.execute_reply.started":"2025-09-12T19:23:20.906118Z","shell.execute_reply":"2025-09-12T19:23:20.921167Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Visualizing audio data\nA random audio file will be used","metadata":{}},{"cell_type":"code","source":"import random\nfor i in range(2):\n    random_index = random.randint(0, df.shape[0])\n    ipd.Audio(df['filepath'].iloc[random_index])\n    audio, sr = load_audio(df['filepath'].iloc[random_index])\n    plt.figure(figsize=(10, 3))\n    pd.Series(audio).plot(figsize=(10, 5),\n                    lw=1,\n                    title=f\"{df['scientific_name'].iloc[random_index]}\",\n                    color=color_pal[0])\n    ## Zoomed in sample to view waves better:\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:20.923100Z","iopub.execute_input":"2025-09-12T19:23:20.923486Z","iopub.status.idle":"2025-09-12T19:23:23.195538Z","shell.execute_reply.started":"2025-09-12T19:23:20.923458Z","shell.execute_reply":"2025-09-12T19:23:23.194519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#### Understanding the audio data\nfor i in range(2):\n    random_index = random.randint(0, df.shape[0])\n    ipd.Audio(df['filepath'].iloc[random_index])\n    audio, sr = load_audio(df['filepath'].iloc[random_index])\n    print(f\"Audio: {audio}\")\n    print(f\"Shape of the audio: {audio.shape}\")\n## The audio file is a numpy array. However, the size of the arrays are different hence we need to pad/trim the arrays to make them the same size\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:23.198010Z","iopub.execute_input":"2025-09-12T19:23:23.198604Z","iopub.status.idle":"2025-09-12T19:23:23.343605Z","shell.execute_reply.started":"2025-09-12T19:23:23.198578Z","shell.execute_reply":"2025-09-12T19:23:23.342604Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Preview a sample of audio spectrograms\nThe STFT represents a signal in the time-frequency domain by computing discrete Fourier transforms (DFT) over short overlapping windows.","metadata":{}},{"cell_type":"code","source":"\"\"\"\n function to preview simple spectrograms in decibels. A decibel is a logarithmic unit that expresses \n the ratio of two values of a physical quantity, often power or intensity.\n \"\"\"\ndef audio_to_spectrogram(audio):\n    D = librosa.stft(audio)\n    S_db = librosa.amplitude_to_db(np.abs(D), ref=np.max)\n    print(S_db.shape)\n\n    fig, ax = plt.subplots(figsize=(10, 5))\n    img = librosa.display.specshow(S_db,\n                                x_axis='time',\n                                y_axis='log',\n                                ax=ax)\n    ax.set_title(f\"{df['scientific_name'].iloc[random_index]} Audio Spectogram\", fontsize=20)\n    fig.colorbar(img, ax=ax, format=f'%0.2f')\n    plt.show()\n\nfor i in range(2):\n    random_index = random.randint(0, df.shape[0])\n    audio, sr = load_audio(df['filepath'].iloc[random_index])\n    audio_to_spectrogram(audio)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:23.344480Z","iopub.execute_input":"2025-09-12T19:23:23.344759Z","iopub.status.idle":"2025-09-12T19:23:25.828826Z","shell.execute_reply.started":"2025-09-12T19:23:23.344737Z","shell.execute_reply":"2025-09-12T19:23:25.828029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nWhile a regular spectrogram uses a linear frequency scale, \na Mel spectrogram uses the Mel scale, which is designed to better reflect how humans perceive sound.\n\"\"\"\ndef audio_to_melspectrogram(audio, sr):\n    S = librosa.feature.melspectrogram(y=audio,\n                                   sr=sr,\n                                   n_mels=128 * 2,)\n    S_db_mel = librosa.amplitude_to_db(S, ref=np.max)\n    print(S_db_mel.shape)\n    fig, ax = plt.subplots(figsize=(10, 5))\n    # Plot the mel spectogram\n    img = librosa.display.specshow(S_db_mel,\n                                x_axis='time',\n                                y_axis='log',\n                                ax=ax)\n    ax.set_title('Mel Spectogram Example', fontsize=20)\n    fig.colorbar(img, ax=ax, format=f'%0.2f')\n    plt.show()\n\nfor i in range(2):\n    random_index = random.randint(0, df.shape[0])\n    audio, sr = load_audio(df['filepath'].iloc[random_index])\n    audio_to_melspectrogram(audio, sr)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:25.829851Z","iopub.execute_input":"2025-09-12T19:23:25.830148Z","iopub.status.idle":"2025-09-12T19:23:26.675856Z","shell.execute_reply.started":"2025-09-12T19:23:25.830126Z","shell.execute_reply":"2025-09-12T19:23:26.674964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\n\nclass AudioUtil:\n    @staticmethod\n    def open(audio_file, target_sample_rate=Config.sample_rate):\n        waveform, sr = torchaudio.load(audio_file)\n        if waveform.shape[0] > 1:\n            waveform = torch.mean(waveform, dim=0, keepdim=True)\n        if sr != target_sample_rate:\n            resampler = T.Resample(sr, target_sample_rate)\n            waveform = resampler(waveform)\n        return waveform, target_sample_rate\n\n    @staticmethod\n    def pad_truncate(waveform, max_len=Config.audio_len):\n        if waveform.shape[1] > max_len:\n            waveform = waveform[:, :max_len]\n        elif waveform.shape[1] < max_len:\n            pad_size = max_len - waveform.shape[1]\n            waveform = F.pad(waveform, (0, pad_size))\n        return waveform\n\n    @staticmethod\n    def spectrogram(waveform, n_fft=Config.nfft, hop_length=Config.hop_length,\n                    n_mels=Config.img_size[0], f_min=Config.fmin, f_max=Config.fmax):\n        mel_spectrogram = T.MelSpectrogram(\n            sample_rate=Config.sample_rate,\n            n_fft=n_fft,\n            hop_length=hop_length,\n            n_mels=n_mels,\n            f_min=f_min,\n            f_max=f_max\n        )(waveform)\n        mel_db = T.AmplitudeToDB()(mel_spectrogram)\n        return mel_db\n\n    @staticmethod\n    def normalize(tensor):\n        min_val, max_val = tensor.min(), tensor.max()\n        tensor = (tensor - min_val) / (max_val - min_val + 1e-6)\n        return tensor\n\n    @staticmethod\n    def to_rgb(tensor):\n        return tensor.repeat(3, 1, 1)\n\n    # ------------------- AUGMENTATIONS -------------------\n    @staticmethod\n    def augment_waveform(waveform, sr):\n        \"\"\"Apply random augmentations directly on the waveform.\"\"\"\n        # Time shifting\n        if random.random() < 0.5:\n            shift = int(random.uniform(-0.1, 0.1) * waveform.shape[1])\n            waveform = torch.roll(waveform, shifts=shift, dims=1)\n\n        # Add Gaussian noise\n        if random.random() < 0.3:\n            noise = torch.randn_like(waveform) * 0.005\n            waveform = waveform + noise\n\n        # Pitch shift (fix: use keyword args for librosa >=0.10)\n        if random.random() < 0.3:\n            n_steps = random.choice([-2, -1, 1, 2])  # semitones\n            waveform_np = waveform.squeeze().numpy()\n            shifted = librosa.effects.pitch_shift(y=waveform_np, sr=sr, n_steps=n_steps)\n            waveform = torch.tensor(shifted).unsqueeze(0)\n\n        return waveform\n\n\n    @staticmethod\n    def augment_spectrogram(spec):\n        \"\"\"Apply random augmentations on the spectrogram (SpecAugment-style).\"\"\"\n        if random.random() < 0.5:\n            spec = T.FrequencyMasking(freq_mask_param=15)(spec)\n        if random.random() < 0.5:\n            spec = T.TimeMasking(time_mask_param=30)(spec)\n        return spec\n\n    # ------------------- MAIN PIPELINE -------------------\n    @staticmethod\n    def process(audio_file, with_label=False, label=None, augment=False):\n        waveform, sr = AudioUtil.open(audio_file)\n        waveform = AudioUtil.pad_truncate(waveform)\n\n        # Apply waveform augmentations\n        if augment:\n            waveform = AudioUtil.augment_waveform(waveform, sr)\n\n        spec = AudioUtil.spectrogram(waveform)\n        \n        # Apply spectrogram augmentations\n        if augment:\n            spec = AudioUtil.augment_spectrogram(spec)\n\n        spec = AudioUtil.normalize(spec)\n        spec = AudioUtil.to_rgb(spec)\n\n        if with_label and label is not None:\n            target = torch.tensor(label, dtype=torch.long)\n            return spec, target\n        else:\n            return spec\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:26.676856Z","iopub.execute_input":"2025-09-12T19:23:26.677180Z","iopub.status.idle":"2025-09-12T19:23:26.691356Z","shell.execute_reply.started":"2025-09-12T19:23:26.677153Z","shell.execute_reply":"2025-09-12T19:23:26.690412Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold\nfrom torch.utils.data import DataLoader, Dataset, Subset\n\nclass BirdClefDataset(Dataset):\n    def __init__(self, df, augment=False):\n        self.df = df.reset_index(drop=True)\n        self.augment = augment\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        spec, label = AudioUtil.process(\n            row.filepath,\n            with_label=True,\n            label=row.target,\n            augment=self.augment\n        )\n        return spec, label\n\ndef balance_classes_for_cv(df, n_splits=5, min_count=3):\n    \"\"\"\n    Hybrid strategy for handling rare classes:\n    - Drop classes with < min_count samples\n    - Oversample classes with count < n_splits\n    - Keep others unchanged\n    \"\"\"\n    dfs = []\n    for label, group in df.groupby(\"target\"):\n        count = len(group)\n        if count < min_count:\n            # Drop ultra-rare classes\n            continue\n        elif count < n_splits:\n            # Oversample to reach n_splits\n            repeat_times = -(-n_splits // count)  # ceil division\n            group = pd.concat([group] * repeat_times, ignore_index=True)\n        dfs.append(group)\n    \n    return pd.concat(dfs, ignore_index=True).reset_index(drop=True)\n\n\ndef build_dataloaders_cv(df, n_splits=5, fold_idx=0, batch_size=32, num_workers=2):\n    # Balance dataset first\n    df_balanced = balance_classes_for_cv(df, n_splits=n_splits)\n\n    # Stratified split\n    skf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=Config.seed)\n    splits = list(skf.split(df_balanced, df_balanced[\"target\"]))\n    train_idx, val_idx = splits[fold_idx]\n    train_df = df_balanced.iloc[train_idx].reset_index(drop=True)\n    val_df = df_balanced.iloc[val_idx].reset_index(drop=True)\n\n    # Build datasets\n    train_ds = BirdClefDataset(train_df, augment=True)\n    val_ds = BirdClefDataset(val_df, augment=False)\n\n    # DataLoaders\n    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=num_workers, pin_memory=True)\n    val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=num_workers, pin_memory=True)\n\n    return train_loader, val_loader, train_df, val_df\n\n\n# def build_dataloaders_cv(df, n_splits=5, batch_size=32, num_workers=2, fold_idx=0):\n#     \"\"\"\n#     Build stratified train/val dataloaders for a given fold.\n    \n#     Args:\n#         df (DataFrame): full metadata dataframe with 'target' column\n#         n_splits (int): number of folds\n#         batch_size (int): batch size\n#         num_workers (int): dataloader workers\n#         fold_idx (int): which fold to use (0 .. n_splits-1)\n#     \"\"\"\n#     skf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=Config.seed)\n\n#     # Split indices\n#     splits = list(skf.split(df, df[\"target\"]))\n#     train_idx, val_idx = splits[fold_idx]\n\n#     train_df = df.iloc[train_idx].reset_index(drop=True)\n#     val_df   = df.iloc[val_idx].reset_index(drop=True)\n\n#     # Datasets\n#     train_dataset = BirdClefDataset(train_df, augment=True)\n#     val_dataset   = BirdClefDataset(val_df, augment=False)\n\n#     # DataLoaders\n#     train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True,\n#                               num_workers=num_workers, pin_memory=True)\n#     val_loader   = DataLoader(val_dataset, batch_size=batch_size, shuffle=False,\n#                               num_workers=num_workers, pin_memory=True)\n\n#     return train_loader, val_loader, train_df, val_df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:26.692501Z","iopub.execute_input":"2025-09-12T19:23:26.692893Z","iopub.status.idle":"2025-09-12T19:23:26.769695Z","shell.execute_reply.started":"2025-09-12T19:23:26.692863Z","shell.execute_reply":"2025-09-12T19:23:26.768986Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def filter_rare_classes(df, n_splits=5):\n    counts = df[\"target\"].value_counts()\n    valid_classes = counts[counts >= n_splits].index\n    return df[df[\"target\"].isin(valid_classes)].reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:26.770847Z","iopub.execute_input":"2025-09-12T19:23:26.771173Z","iopub.status.idle":"2025-09-12T19:23:26.776075Z","shell.execute_reply.started":"2025-09-12T19:23:26.771144Z","shell.execute_reply":"2025-09-12T19:23:26.775129Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example: 5-fold cross validation\nfor fold in range(5):\n    print(f\"\\n===== Fold {fold+1} =====\")\n    train_loader, val_loader, train_df, val_df = build_dataloaders_cv(df, n_splits=5, fold_idx=0)\n    print(\"Train size:\", len(train_df), \" Val size:\", len(val_df))\n\n    # Inspect one batch\n    specs, labels = next(iter(train_loader))\n    print(\"Batch shape:\", specs.shape, \" Labels shape:\", labels.shape)\n\n    # You can now train your model on this fold\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:23:26.776967Z","iopub.execute_input":"2025-09-12T19:23:26.777230Z","iopub.status.idle":"2025-09-12T19:23:50.139687Z","shell.execute_reply.started":"2025-09-12T19:23:26.777210Z","shell.execute_reply":"2025-09-12T19:23:50.138461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import timm\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.cuda.amp import autocast, GradScaler\nfrom tqdm import tqdm\n\n# ------------------- Model -------------------\n# def create_model(num_classes=Config.num_classes, pretrained=True):\n#     model = timm.create_model('efficientnetv2', pretrained=pretrained)\n#     in_features = model.classifier.in_features\n#     model.classifier = nn.Linear(in_features, num_classes)\n#     return model.to(Config.device)\n# Choose one of these specific EfficientNetV2 variants:\ndef create_model(num_classes=Config.num_classes, pretrained=True):\n    # Common EfficientNetV2 variants:\n    # 'efficientnetv2_rw_t' - Tiny\n    # 'efficientnetv2_rw_s' - Small  \n    # 'efficientnetv2_rw_m' - Medium\n    # 'efficientnetv2_rw_l' - Large\n    \n    model = timm.create_model('efficientnetv2_rw_s', pretrained=pretrained)  # Example: Small variant\n    in_features = model.classifier.in_features\n    model.classifier = nn.Linear(in_features, num_classes)\n    return model.to(Config.device)\n\n# ------------------- Optimizer + Scheduler -------------------\ndef create_optimizer(model, lr=1e-4, weight_decay=1e-5):\n    optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)\n    return optimizer\n\ndef create_scheduler(optimizer, train_loader, epochs=10):\n    scheduler = optim.lr_scheduler.OneCycleLR(\n        optimizer,\n        max_lr=1e-3,\n        steps_per_epoch=len(train_loader),\n        epochs=epochs,\n        pct_start=0.1,\n        anneal_strategy='cos',\n        div_factor=10,\n        final_div_factor=100,\n        three_phase=False\n    )\n    return scheduler\n\n# ------------------- Training Loop -------------------\ndef train_one_epoch(model, train_loader, criterion, optimizer, scaler):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    for specs, labels in tqdm(train_loader, desc=\"Training\", leave=False):\n        specs, labels = specs.to(Config.device), labels.to(Config.device)\n\n        optimizer.zero_grad()\n        with autocast():\n            outputs = model(specs)\n            loss = criterion(outputs, labels)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item() * specs.size(0)\n        _, preds = torch.max(outputs, 1)\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n\n    epoch_loss = running_loss / total\n    epoch_acc = correct / total\n    return epoch_loss, epoch_acc\n\ndef validate(model, val_loader, criterion):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    with torch.no_grad():\n        for specs, labels in tqdm(val_loader, desc=\"Validation\", leave=False):\n            specs, labels = specs.to(Config.device), labels.to(Config.device)\n            outputs = model(specs)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * specs.size(0)\n            _, preds = torch.max(outputs, 1)\n            correct += (preds == labels).sum().item()\n            total += labels.size(0)\n\n    epoch_loss = running_loss / total\n    epoch_acc = correct / total\n    return epoch_loss, epoch_acc\n\n# ------------------- Full Training -------------------\ndef train_model(model, train_loader, val_loader, epochs=10, lr=1e-4):\n    criterion = nn.CrossEntropyLoss()\n    optimizer = create_optimizer(model, lr)\n    scheduler = create_scheduler(optimizer, train_loader, epochs)\n    scaler = GradScaler()\n\n    best_val_acc = 0.0\n\n    for epoch in range(epochs):\n        print(f\"\\nEpoch {epoch+1}/{epochs}\")\n\n        train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, scaler)\n        val_loss, val_acc = validate(model, val_loader, criterion)\n\n        scheduler.step()\n\n        print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n        print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f}\")\n\n        # Save best model\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            torch.save(model.state_dict(), \"best_model.pth\")\n            print(\"Saved Best Model ✅\")\n\n    print(\"Training complete. Best Val Acc:\", best_val_acc)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:39:23.692513Z","iopub.execute_input":"2025-09-12T19:39:23.693547Z","iopub.status.idle":"2025-09-12T19:39:23.709558Z","shell.execute_reply.started":"2025-09-12T19:39:23.693505Z","shell.execute_reply":"2025-09-12T19:39:23.708673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create model\nmodel = create_model()\n\n# Build DataLoaders (example fold 0)\ntrain_loader, val_loader, train_df, val_df = build_dataloaders_cv(df, n_splits=5, fold_idx=0, batch_size=32)\n\n# Train\ntrain_model(model, train_loader, val_loader, epochs=10, lr=1e-4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T19:39:28.205434Z","iopub.execute_input":"2025-09-12T19:39:28.206196Z","execution_failed":"2025-09-12T19:39:48.697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}