{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":11295156,"sourceType":"datasetVersion","datasetId":7062748},{"sourceId":11386809,"sourceType":"datasetVersion","datasetId":7130356},{"sourceId":3732,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":2659,"modelId":312}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# importing libraries\nimport librosa\nimport numpy as np\nimport librosa.display\nimport IPython.display as ipd\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport torch.nn as nn\nimport torch.optim as optim\nimport glob\nimport torch\nimport ast\n\nimport random\nfrom torchvision import models\nfrom sklearn.model_selection import StratifiedKFold\n\nfrom warnings import filterwarnings\nfilterwarnings('ignore')\n\nimport seaborn as sns\nfrom ast import literal_eval\nimport os\nfrom tqdm import tqdm\nfrom sklearn.model_selection import KFold, StratifiedKFold\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nimport torch.nn.functional as F\nimport gc\nfrom torchvision import transforms\nimport geopandas as gpd\nfrom shapely.geometry import Point\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\nimport joblib\npd.set_option('display.max_colwidth', None)\nfrom concurrent.futures import ProcessPoolExecutor, as_completed\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.346744Z","iopub.execute_input":"2025-04-16T21:35:06.347134Z","iopub.status.idle":"2025-04-16T21:35:06.353457Z","shell.execute_reply.started":"2025-04-16T21:35:06.347104Z","shell.execute_reply":"2025-04-16T21:35:06.352665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# create the config class\nclass config:\n    train_audio = '/kaggle/input/birdclef-2025/train_audio'\n    train_csv = '/kaggle/input/birdclef-2025/train.csv'\n    sample_solution = '/kaggle/input/birdclef-2025/sample_submission.csv'\n    test_soundscape = '/kaggle/input/birdclef-2025/test_soundscapes'\n    save_path = '/kaggle/working/'\n\n    sampling_rate = 32000\n    num_classes = 206\n    n_mels = 128\n    fmin = 40\n    fmax = 15000\n    chunk_length = 10  # seconds\n    n_fft = 1024\n    hop_length = 512\n    seed = 42\n\n    batch_size = 64\n    epochs = 12\n    learning_rate = 2e-3\n    num_folds = 5\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.354572Z","iopub.execute_input":"2025-04-16T21:35:06.354882Z","iopub.status.idle":"2025-04-16T21:35:06.373627Z","shell.execute_reply.started":"2025-04-16T21:35:06.354860Z","shell.execute_reply":"2025-04-16T21:35:06.372954Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def set_seed(seed):\n    # see in config see I don't want to return anything cool\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    \n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n    \n    print(f\"All done! Set seed: {seed}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.375341Z","iopub.execute_input":"2025-04-16T21:35:06.375543Z","iopub.status.idle":"2025-04-16T21:35:06.386340Z","shell.execute_reply.started":"2025-04-16T21:35:06.375526Z","shell.execute_reply":"2025-04-16T21:35:06.385539Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Working on train.csv file","metadata":{}},{"cell_type":"code","source":"# Load the train dataframe\ntrain_df = pd.read_csv(config.train_csv)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.387454Z","iopub.execute_input":"2025-04-16T21:35:06.387711Z","iopub.status.idle":"2025-04-16T21:35:06.492785Z","shell.execute_reply.started":"2025-04-16T21:35:06.387683Z","shell.execute_reply":"2025-04-16T21:35:06.492107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df.dtypes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.493516Z","iopub.execute_input":"2025-04-16T21:35:06.493811Z","iopub.status.idle":"2025-04-16T21:35:06.499497Z","shell.execute_reply.started":"2025-04-16T21:35:06.493778Z","shell.execute_reply":"2025-04-16T21:35:06.498745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# deal with secondary_label , type and filenames columns\nfor col in ('secondary_labels', 'type'):\n    train_df[col] = train_df[col].apply(lambda x: \"###\".join(literal_eval(x)))\n\ntrain_df['filename'] = train_df['filename'].apply(lambda x: config.train_audio + \"/\" + x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.500309Z","iopub.execute_input":"2025-04-16T21:35:06.500545Z","iopub.status.idle":"2025-04-16T21:35:06.862373Z","shell.execute_reply.started":"2025-04-16T21:35:06.500519Z","shell.execute_reply":"2025-04-16T21:35:06.861760Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df.sample(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.863152Z","iopub.execute_input":"2025-04-16T21:35:06.863440Z","iopub.status.idle":"2025-04-16T21:35:06.879444Z","shell.execute_reply.started":"2025-04-16T21:35:06.863419Z","shell.execute_reply":"2025-04-16T21:35:06.878717Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_df['filename'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.880277Z","iopub.execute_input":"2025-04-16T21:35:06.880528Z","iopub.status.idle":"2025-04-16T21:35:06.893834Z","shell.execute_reply.started":"2025-04-16T21:35:06.880493Z","shell.execute_reply":"2025-04-16T21:35:06.893169Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# checking duplicate sound file \ntrain_df['filename'].duplicated(keep=False).sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.896371Z","iopub.execute_input":"2025-04-16T21:35:06.896578Z","iopub.status.idle":"2025-04-16T21:35:06.912075Z","shell.execute_reply.started":"2025-04-16T21:35:06.896560Z","shell.execute_reply":"2025-04-16T21:35:06.911265Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Plot the distribution of classes in world map","metadata":{}},{"cell_type":"code","source":"# plot class locations using latitude and longitude columns\n# geometry = [Point(lon, lat) for lon, lat in zip(train_df['longitude'], train_df['latitude'])]\n# gdf = gpd.GeoDataFrame(train_df, geometry=geometry)\n\n\n# fig, ax = plt.subplots(figsize=(30, 20)) \n\n\n# gdf.plot(ax=ax, marker='o', color='red', markersize=30)\n\n\n# world = gpd.read_file(gpd.datasets.get_path('naturalearth_lowres'))\n\n\n# world.boundary.plot(ax=ax, color='black', linewidth=2)\n\n# ax.set_title('Locations on Map with Borders')\n# ax.set_xlabel('Longitude')\n# ax.set_ylabel('Latitude')\n\n\n# plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.913148Z","iopub.execute_input":"2025-04-16T21:35:06.913381Z","iopub.status.idle":"2025-04-16T21:35:06.926574Z","shell.execute_reply.started":"2025-04-16T21:35:06.913362Z","shell.execute_reply":"2025-04-16T21:35:06.925844Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Plot sound ratings ","metadata":{}},{"cell_type":"code","source":"# Disrtibution of sound quality ratings\n# plt.style.use('seaborn')\n\n\n# plt.figure(figsize=(10, 6))\n\n\n# sns.countplot(data=train_df, x='rating')\n\n\n# plt.title('Distribution of Sound Quality Ratings', fontsize=14, pad=15)\n# plt.xlabel('Rating', fontsize=12)\n# plt.ylabel('Count', fontsize=12)\n\n\n# for i in plt.gca().containers:\n#     plt.gca().bar_label(i)\n\n\n# plt.tight_layout()\n\n\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.927444Z","iopub.execute_input":"2025-04-16T21:35:06.927699Z","iopub.status.idle":"2025-04-16T21:35:06.941460Z","shell.execute_reply.started":"2025-04-16T21:35:06.927679Z","shell.execute_reply":"2025-04-16T21:35:06.940661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_df.isnull().sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.942159Z","iopub.execute_input":"2025-04-16T21:35:06.942366Z","iopub.status.idle":"2025-04-16T21:35:06.954980Z","shell.execute_reply.started":"2025-04-16T21:35:06.942349Z","shell.execute_reply":"2025-04-16T21:35:06.954292Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Plot distribution of the classes","metadata":{}},{"cell_type":"code","source":"# plt.figure(figsize=(20, 9))\n# sns.countplot(data=train_df, x='primary_label')\n# plt.xticks(rotation=90, ha='right')\n# plt.title('Distribution of Sound Classes', fontsize=14, pad=15)\n# plt.xlabel('Class', fontsize=12)\n# plt.ylabel('Count', fontsize=12)\n# plt.tight_layout()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.955659Z","iopub.execute_input":"2025-04-16T21:35:06.955929Z","iopub.status.idle":"2025-04-16T21:35:06.971416Z","shell.execute_reply.started":"2025-04-16T21:35:06.955902Z","shell.execute_reply":"2025-04-16T21:35:06.970840Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Top 10 most frequent classes","metadata":{}},{"cell_type":"code","source":"# Get top 10 most frequent classes\n# top_10_classes = train_df['primary_label'].value_counts().head(10)\n\n\n# plt.figure(figsize=(12, 6))\n# sns.barplot(x=top_10_classes.values, y=top_10_classes.index)\n# plt.title('Top 10 Most Frequent Classes', fontsize=14, pad=15)\n# plt.xlabel('Count', fontsize=12)\n# plt.ylabel('Class', fontsize=12)\n\n\n# for i, v in enumerate(top_10_classes.values):\n#     plt.text(v, i, str(v), va='center')\n\n# plt.tight_layout()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.972057Z","iopub.execute_input":"2025-04-16T21:35:06.972335Z","iopub.status.idle":"2025-04-16T21:35:06.990772Z","shell.execute_reply.started":"2025-04-16T21:35:06.972303Z","shell.execute_reply":"2025-04-16T21:35:06.990121Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# converting audio to mel spectrogram","metadata":{}},{"cell_type":"code","source":"def load_audio(path, duration=config.chunk_length, sr=config.sampling_rate):\n    y, _ = librosa.load(path, sr=sr, mono=True)\n    length = sr * duration\n\n    if len(y) < length:\n        y = np.pad(y, (0, length - len(y)))\n    else:\n        y = y[:length]\n\n    return y\n\ndef audio_to_melspec(audio):\n    melspec = librosa.feature.melspectrogram(\n        y=audio,\n        sr=config.sampling_rate,\n        n_fft=config.n_fft,\n        hop_length=config.hop_length,\n        n_mels=config.n_mels,\n        fmin=config.fmin,\n        fmax=config.fmax\n    )\n    melspec_db = librosa.power_to_db(melspec, ref=np.max)\n    return melspec_db","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:06.991487Z","iopub.execute_input":"2025-04-16T21:35:06.991778Z","iopub.status.idle":"2025-04-16T21:35:07.005271Z","shell.execute_reply.started":"2025-04-16T21:35:06.991747Z","shell.execute_reply":"2025-04-16T21:35:07.004432Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# BirdDataset","metadata":{}},{"cell_type":"code","source":"class BirdDataset(Dataset):\n    def __init__(self, df, label_map, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n        self.label_map = label_map\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        path = os.path.join(config.train_audio, row['filename'])\n        audio = load_audio(path)\n        melspec = audio_to_melspec(audio)\n\n        # Normalize to fixed shape\n        expected_shape = (config.n_mels, int(config.chunk_length * config.sampling_rate / config.hop_length))\n        if melspec.shape[1] < expected_shape[1]:\n            pad_width = expected_shape[1] - melspec.shape[1]\n            melspec = np.pad(melspec, ((0, 0), (0, pad_width)))\n        else:\n            melspec = melspec[:, :expected_shape[1]]\n\n        # Convert to 3-channel\n        melspec_rgb = np.stack([melspec, melspec, melspec])\n        melspec_tensor = torch.tensor(melspec_rgb).float()\n\n        if self.transform:\n            melspec_tensor = self.transform(melspec_tensor)\n\n        label = torch.zeros(config.num_classes)\n        label[row['primary_label']] = 1.0\n\n        if isinstance(row['secondary_labels'], str):\n            for sec in row['secondary_labels'].split(\"###\"):\n                if sec in self.label_map:\n                    label[self.label_map[sec]] = 0.5\n\n        return melspec_tensor, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:07.006071Z","iopub.execute_input":"2025-04-16T21:35:07.006284Z","iopub.status.idle":"2025-04-16T21:35:07.016651Z","shell.execute_reply.started":"2025-04-16T21:35:07.006267Z","shell.execute_reply":"2025-04-16T21:35:07.015971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_transforms():\n    return transforms.Compose([\n        transforms.RandomHorizontalFlip(),\n        transforms.RandomErasing(p=0.5),\n        transforms.Normalize([0.485]*3, [0.229]*3)\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:07.017363Z","iopub.execute_input":"2025-04-16T21:35:07.017652Z","iopub.status.idle":"2025-04-16T21:35:07.033246Z","shell.execute_reply.started":"2025-04-16T21:35:07.017624Z","shell.execute_reply":"2025-04-16T21:35:07.032553Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# EfficientNet B0 model\n","metadata":{}},{"cell_type":"code","source":"def get_model():\n    model = models.efficientnet_b3(pretrained=True)\n    model.classifier[1] = nn.Linear(model.classifier[1].in_features, config.num_classes)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:07.033911Z","iopub.execute_input":"2025-04-16T21:35:07.034150Z","iopub.status.idle":"2025-04-16T21:35:07.045894Z","shell.execute_reply.started":"2025-04-16T21:35:07.034131Z","shell.execute_reply":"2025-04-16T21:35:07.045091Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train Epoch","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, dataloader, criterion, optimizer, scaler):\n    model.train()\n    total_loss = 0\n\n    for inputs, targets in tqdm(dataloader, desc=\"Training\"):\n        inputs, targets = inputs.to(config.device), targets.to(config.device)\n\n        optimizer.zero_grad()\n        with torch.cuda.amp.autocast():\n            outputs = model(inputs)\n            loss = criterion(outputs, targets)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        total_loss += loss.item()\n\n    return total_loss / len(dataloader)\n\n# def train_epoch(model, dataloader, criterion, optimizer):\n#     model.train()\n#     total_loss = 0\n\n#     for inputs, targets in tqdm(dataloader, desc=\"Training\"):\n#         inputs, targets = inputs.to(config.device), targets.to(config.device)\n\n#         optimizer.zero_grad()\n#         outputs = model(inputs)\n#         loss = criterion(outputs, targets)\n#         loss.backward()\n#         optimizer.step()\n\n#         total_loss += loss.item()\n\n#     return total_loss / len(dataloader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:07.046800Z","iopub.execute_input":"2025-04-16T21:35:07.047079Z","iopub.status.idle":"2025-04-16T21:35:07.059542Z","shell.execute_reply.started":"2025-04-16T21:35:07.047051Z","shell.execute_reply":"2025-04-16T21:35:07.058756Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# validate epoch","metadata":{}},{"cell_type":"code","source":"def validate_epoch(model, dataloader, criterion):\n    model.eval()\n    total_loss = 0\n\n    with torch.no_grad():\n        for inputs, targets in tqdm(dataloader, desc=\"Validating\"):\n            inputs, targets = inputs.to(config.device), targets.to(config.device)\n            with torch.cuda.amp.autocast():\n                outputs = model(inputs)\n                loss = criterion(outputs, targets)\n            total_loss += loss.item()\n\n    return total_loss / len(dataloader)\n\n\n# def validate_epoch(model, dataloader, criterion):\n#     model.eval()\n#     total_loss = 0\n\n#     with torch.no_grad():\n#         for inputs, targets in tqdm(dataloader, desc=\"Validating\"):\n#             inputs, targets = inputs.to(config.device), targets.to(config.device)\n#             outputs = model(inputs)\n#             loss = criterion(outputs, targets)\n#             total_loss += loss.item()\n\n#     return total_loss / len(dataloader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:07.060370Z","iopub.execute_input":"2025-04-16T21:35:07.060615Z","iopub.status.idle":"2025-04-16T21:35:07.075809Z","shell.execute_reply.started":"2025-04-16T21:35:07.060574Z","shell.execute_reply":"2025-04-16T21:35:07.075138Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train model","metadata":{}},{"cell_type":"code","source":"# def run_training(df):\n#     label_map = {label: i for i, label in enumerate(sorted(df['primary_label'].unique()))}\n#     df['primary_label'] = df['primary_label'].map(label_map)\n\n#     skf = StratifiedKFold(n_splits=config.num_folds, shuffle=True, random_state=config.seed)\n\n#     for fold, (train_idx, val_idx) in enumerate(skf.split(df, df['primary_label'])):\n#         if fold != 1:\n#             continue  # Skip all folds except fold 1\n\n#         print(f\"\\n=== Fold {fold+1} ===\")\n#         train_df = df.iloc[train_idx]\n#         val_df = df.iloc[val_idx]\n\n#         train_dataset = BirdDataset(train_df, label_map, transform=get_transforms())\n#         val_dataset = BirdDataset(val_df, label_map)\n\n#         train_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True, num_workers=2)\n#         val_loader = DataLoader(val_dataset, batch_size=config.batch_size, shuffle=False, num_workers=2)\n\n#         model = get_model().to(config.device)\n#         criterion = nn.BCEWithLogitsLoss()\n#         optimizer = optim.AdamW(model.parameters(), lr=config.learning_rate)\n\n#         best_loss = float('inf')\n#         for epoch in range(config.epochs):\n#             print(f\"\\nEpoch {epoch+1}/{config.epochs}\")\n#             train_loss = train_epoch(model, train_loader, criterion, optimizer)\n#             val_loss = validate_epoch(model, val_loader, criterion)\n\n#             print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n\n#             if val_loss < best_loss:\n#                 best_loss = val_loss\n#                 torch.save(model.state_dict(), f\"{config.save_path}/best_fold{fold}.pth\")\n#                 print(\"Model saved!\")\n\n#         break  # Optional: stops after fold 1 (since others are skipped anyway)\ndef run_training(df):\n    label_map = {label: i for i, label in enumerate(sorted(df['primary_label'].unique()))}\n    df['primary_label'] = df['primary_label'].map(label_map)\n\n    skf = StratifiedKFold(n_splits=config.num_folds, shuffle=True, random_state=config.seed)\n\n    for fold, (train_idx, val_idx) in enumerate(skf.split(df, df['primary_label'])):\n        if fold != 1:\n            continue  # Skip all folds except fold 1\n\n        print(f\"\\n=== Fold {fold+1} ===\")\n        train_df = df.iloc[train_idx]\n        val_df = df.iloc[val_idx]\n\n        train_dataset = BirdDataset(train_df, label_map, transform=get_transforms())\n        val_dataset = BirdDataset(val_df, label_map)\n\n        train_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True, num_workers=1, pin_memory=False)\n        val_loader = DataLoader(val_dataset, batch_size=config.batch_size, shuffle=False, num_workers=1, pin_memory=False)\n\n        model = get_model().to(config.device)\n        criterion = nn.CrossEntropyLoss()\n        optimizer = optim.AdamW(model.parameters(), lr=config.learning_rate)\n        scaler = torch.cuda.amp.GradScaler()\n\n        best_loss = float('inf')\n        for epoch in range(config.epochs):\n            print(f\"\\nEpoch {epoch+1}/{config.epochs}\")\n            model.train()\n            total_train_loss = 0\n\n            for inputs, targets in tqdm(train_loader, desc=\"Training\"):\n                inputs, targets = inputs.to(config.device), targets.to(config.device)\n\n                optimizer.zero_grad()\n                with torch.cuda.amp.autocast():\n                    outputs = model(inputs)\n                    loss = criterion(outputs, targets)\n\n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n\n                total_train_loss += loss.item()\n\n            train_loss = total_train_loss / len(train_loader)\n\n            # Validation\n            model.eval()\n            total_val_loss = 0\n            with torch.no_grad():\n                for inputs, targets in tqdm(val_loader, desc=\"Validating\"):\n                    inputs, targets = inputs.to(config.device), targets.to(config.device)\n                    with torch.cuda.amp.autocast():\n                        outputs = model(inputs)\n                        loss = criterion(outputs, targets)\n                    total_val_loss += loss.item()\n\n            val_loss = total_val_loss / len(val_loader)\n\n            print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n\n            if val_loss < best_loss:\n                best_loss = val_loss\n                torch.save(model.state_dict(), f\"{config.save_path}/best_fold{fold}_efficientnet_b3_new.pth\")\n                print(\"Model saved!\")\n\n            torch.cuda.empty_cache()\n            gc.collect()\n\n        break  # Optional: stops after fold 1\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:07.076455Z","iopub.execute_input":"2025-04-16T21:35:07.076684Z","iopub.status.idle":"2025-04-16T21:35:07.087101Z","shell.execute_reply.started":"2025-04-16T21:35:07.076665Z","shell.execute_reply":"2025-04-16T21:35:07.086276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_training(train_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T21:35:07.088038Z","iopub.execute_input":"2025-04-16T21:35:07.088337Z","execution_failed":"2025-04-16T22:04:27.301Z"}},"outputs":[],"execution_count":null}]}