{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":91844,"databundleVersionId":11361821},{"sourceType":"kernelVersion","sourceId":239177129}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Import necessary libraries\nimport os\nimport pandas as pd\nimport numpy as np\nimport librosa\nimport librosa.display\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nfrom tqdm.notebook import tqdm # Or standard tqdm if not in notebook\nimport time\nimport copy # To save best model state\nimport random\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.models as models\nimport torchaudio.transforms as T # For SpecAugment","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:29:45.944737Z","iopub.execute_input":"2025-05-08T12:29:45.945043Z","iopub.status.idle":"2025-05-08T12:29:45.950131Z","shell.execute_reply.started":"2025-05-08T12:29:45.945019Z","shell.execute_reply":"2025-05-08T12:29:45.949352Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 1. Configuration & Load Data ---\n# Adjust these paths based on your Kaggle environment\n# Create dummy paths and data if not on Kaggle for local testing\nIS_KAGGLE = os.path.exists(\"/kaggle/input\")\nif IS_KAGGLE:\n    CSV_PATH = \"/kaggle/input/birdclef-2025/train.csv\"\n    AUDIO_BASE_PATH = \"/kaggle/input/birdclef-2025/train_audio/\"\nelse:\n    print(\"Kaggle paths not found. Creating dummy data for local testing.\")\n    if not os.path.exists(\"./kaggle/input/birdclef-2025/\"):\n        os.makedirs(\"./kaggle/input/birdclef-2025/\", exist_ok=True)\n    if not os.path.exists(\"./kaggle/input/birdclef-2025/train_audio/\"):\n        os.makedirs(\"./kaggle/input/birdclef-2025/train_audio/\", exist_ok=True)\n\n    CSV_PATH = \"./kaggle/input/birdclef-2025/train.csv\"\n    AUDIO_BASE_PATH = \"./kaggle/input/birdclef-2025/train_audio/\"\n\n    # Create dummy CSV\n    dummy_filenames = [f'bird_{i}.ogg' for i in range(20)] # Increased dummy files\n    dummy_data = {'filename': dummy_filenames,\n                  'primary_label': [f'species_{i%5}' for i in range(20)], # 5 dummy species\n                  'secondary_labels': [[] for _ in range(20)]}\n    dummy_df = pd.DataFrame(dummy_data)\n    dummy_df.to_csv(CSV_PATH, index=False)\n\n    # Create dummy audio files\n    try:\n        import soundfile as sf\n        for fname in dummy_filenames:\n            dummy_audio = np.random.uniform(-0.5, 0.5, 32000 * 5) # 5 seconds of random noise\n            sf.write(os.path.join(AUDIO_BASE_PATH, fname), dummy_audio, 32000)\n        print(f\"Dummy data created at {CSV_PATH} and {AUDIO_BASE_PATH}\")\n    except ImportError:\n        print(\"Please install 'soundfile' to create dummy audio files: pip install soundfile\")\n        exit()\n\n\n# Audio Processing Parameters\nSAMPLE_RATE = 32000\nDURATION = 5  # seconds\nN_MELS = 128\nFMIN = 20\nFMAX = 14000 # Birds typically don't vocalize much higher than this\nHOP_LENGTH = int(SAMPLE_RATE * 0.01)  # 10ms hop: 32000 * 0.01 = 320\nN_FFT = int(SAMPLE_RATE * 0.025)     # 25ms window: 32000 * 0.025 = 800, round to 1024 for FFT efficiency\nN_FFT = 1024 # Or 2048 as you had, both are fine. Let's try 1024.\n\n# Augmentation Parameters\nAPPLY_AUGMENTATION_PROB = 0.6 # Probability of applying augmentations to a training sample\nNOISE_LEVEL = 0.005\nTIME_SHIFT_FRACTION = 0.2 # Max fraction of duration to shift\nPITCH_SHIFT_STEPS = 2 # Max semitones for pitch shift\n\n# Model & Training Parameters\nMODEL_NAME = \"efficientnet_b3\" # e.g., \"efficientnet_b0\", \"efficientnet_b2\"\nBATCH_SIZE = 32\nEPOCHS = 25 # Pre-trained models can benefit from more epochs with good augmentation\nLEARNING_RATE = 3e-4 # Often a good starting LR for AdamW with pre-trained models\nWEIGHT_DECAY = 0.01 # For AdamW\nVALIDATION_SPLIT = 0.2\nRANDOM_SEED = 42\nNUM_WORKERS = 2 # os.cpu_count() // 2 if you have many cores\nLABEL_SMOOTHING = 0.1 # Set to 0 to disable\nPRELOAD_DATA_IN_RAM = True # Set to True if you have enough RAM and want faster epochs after initial load\n\n# --- Set Seed for Reproducibility ---\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed) # if you are using multi-GPU.\n    # 영향 줄 수 있는 설정들\n    # torch.backends.cudnn.deterministic = True # Can slow down training\n    # torch.backends.cudnn.benchmark = False    # Can slow down training if input sizes vary\n\nseed_everything(RANDOM_SEED)\n\n# --- Determine Device ---\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# --- Load the training metadata ---\ntry:\n    data = pd.read_csv(CSV_PATH)\n    print(f\"Successfully loaded {CSV_PATH}\")\n    print(f\"Data shape: {data.shape}\")\nexcept FileNotFoundError:\n    print(f\"Error: Could not find {CSV_PATH}\")\n    exit()\n\n# Construct full audio file paths\ndata['full_path'] = data['filename'].apply(lambda x: os.path.join(AUDIO_BASE_PATH, x))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:36:44.358237Z","iopub.execute_input":"2025-05-08T12:36:44.358944Z","iopub.status.idle":"2025-05-08T12:36:44.543393Z","shell.execute_reply.started":"2025-05-08T12:36:44.358920Z","shell.execute_reply":"2025-05-08T12:36:44.542631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 2. Label Encoding ---\nencoder = LabelEncoder()\ndata['label_encoded'] = encoder.fit_transform(data['primary_label'])\nNUM_CLASSES = len(encoder.classes_)\nprint(f\"\\nNumber of unique classes: {NUM_CLASSES}\")\n# label_to_int = {label: i for i, label in enumerate(encoder.classes_)}\n# int_to_label = {i: label for i, label in enumerate(encoder.classes_)}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:32:17.554833Z","iopub.execute_input":"2025-05-08T12:32:17.555504Z","iopub.status.idle":"2025-05-08T12:32:17.565182Z","shell.execute_reply.started":"2025-05-08T12:32:17.555478Z","shell.execute_reply":"2025-05-08T12:32:17.564649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 3. Train/Validation Split ---\ntrain_df, val_df = train_test_split(\n    data,\n    test_size=VALIDATION_SPLIT,\n    random_state=RANDOM_SEED,\n    stratify=data['label_encoded'] # Crucial for imbalanced datasets\n)\nprint(f\"\\nTraining set size: {len(train_df)}\")\nprint(f\"Validation set size: {len(val_df)}\")\n\n# Calculate expected spectrogram shape (for reference and padding)\n# For librosa.stft, frame length is N_FFT. Number of frames is roughly len(y) / hop_length.\n# For melspectrogram, the time dimension is ceil(samples / hop_length)\n# After padding/truncating audio to DURATION * SAMPLE_RATE:\nTARGET_LENGTH_SAMPLES = int(SAMPLE_RATE * DURATION)\nN_FRAMES = int(np.ceil(TARGET_LENGTH_SAMPLES / HOP_LENGTH)) + 1 # Librosa can add a frame due to centering\nprint(f\"Expected Spectrogram Shape (H, W) after processing: ({N_MELS}, {N_FRAMES})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:32:18.925871Z","iopub.execute_input":"2025-05-08T12:32:18.926456Z","iopub.status.idle":"2025-05-08T12:32:18.959794Z","shell.execute_reply.started":"2025-05-08T12:32:18.926433Z","shell.execute_reply":"2025-05-08T12:32:18.959091Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 4. Audio Preprocessing & Augmentation Function ---\ndef load_and_preprocess_audio(\n    file_path, target_sr=SAMPLE_RATE, duration_samples=TARGET_LENGTH_SAMPLES,\n    n_mels=N_MELS, hop_length=HOP_LENGTH, n_fft=N_FFT,\n    fmin=FMIN, fmax=FMAX, is_train=False):\n    \"\"\"Loads, augments (if train), pads/truncates, and computes Mel spectrogram.\"\"\"\n    try:\n        y, sr = librosa.load(file_path, sr=None) # Load native SR\n        if sr != target_sr:\n            y = librosa.resample(y, orig_sr=sr, target_sr=target_sr)\n\n        # --- Augmentations (on raw audio, only for training) ---\n        if is_train and random.random() < APPLY_AUGMENTATION_PROB:\n            # 1. Time Shifting\n            if random.random() < 0.5:\n                max_shift = int(len(y) * TIME_SHIFT_FRACTION)\n                shift = random.randint(-max_shift, max_shift)\n                y = np.roll(y, shift)\n\n            # 2. Adding Noise\n            if random.random() < 0.5:\n                noise = np.random.randn(len(y)) * NOISE_LEVEL\n                y = y + noise\n\n            # 3. Pitch Shifting (can be slow, apply with lower probability if needed)\n            # if random.random() < 0.3:\n            #     n_steps = random.uniform(-PITCH_SHIFT_STEPS, PITCH_SHIFT_STEPS)\n            #     y = librosa.effects.pitch_shift(y, sr=target_sr, n_steps=n_steps)\n\n        # Pad or truncate audio to target_length_samples\n        if len(y) < duration_samples:\n            padding = duration_samples - len(y)\n            offset = random.randint(0, padding) if is_train else padding // 2 # Pad randomly for train\n            y = np.pad(y, (offset, padding - offset), mode='constant')\n        elif len(y) > duration_samples:\n            start = random.randint(0, len(y) - duration_samples) if is_train else 0 # Random crop for train\n            y = y[start : start + duration_samples]\n\n        mel_spec = librosa.feature.melspectrogram(\n            y=y, sr=target_sr, n_fft=n_fft, hop_length=hop_length,\n            n_mels=n_mels, fmin=fmin, fmax=fmax, window='hann' # Hann window is common\n        )\n        mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max)\n\n        # Normalize to [0, 1] (or standardize if preferred)\n        min_val = np.min(mel_spec_db)\n        max_val = np.max(mel_spec_db)\n        if max_val > min_val:\n            mel_spec_db = (mel_spec_db - min_val) / (max_val - min_val)\n        else: # Handle silent clips\n            mel_spec_db = np.zeros_like(mel_spec_db)\n\n        # Ensure consistent shape (especially time dimension) due to potential minor librosa variations\n        if mel_spec_db.shape[1] < N_FRAMES:\n            pad_width = N_FRAMES - mel_spec_db.shape[1]\n            mel_spec_db = np.pad(mel_spec_db, ((0,0), (0, pad_width)), mode='constant', constant_values=0.)\n        elif mel_spec_db.shape[1] > N_FRAMES:\n            mel_spec_db = mel_spec_db[:, :N_FRAMES]\n\n\n        return mel_spec_db\n    except Exception as e:\n        print(f\"Error processing {file_path}: {e}\")\n        return np.zeros((n_mels, N_FRAMES)) # Return zeros on error, with expected shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:32:20.331704Z","iopub.execute_input":"2025-05-08T12:32:20.332249Z","iopub.status.idle":"2025-05-08T12:32:20.341920Z","shell.execute_reply.started":"2025-05-08T12:32:20.332228Z","shell.execute_reply":"2025-05-08T12:32:20.341112Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 5. Create PyTorch Dataset ---\nclass BirdSoundDataset(Dataset):\n    def __init__(self, dataframe, is_train=False, preload_in_ram=PRELOAD_DATA_IN_RAM,\n                 apply_spec_augment=True):\n        self.dataframe = dataframe\n        self.is_train = is_train\n        self.preload_in_ram = preload_in_ram\n        self.apply_spec_augment = apply_spec_augment\n\n        if self.is_train and self.apply_spec_augment:\n            # SpecAugment: Applied on the Mel spectrogram tensor\n            self.spec_augmenter =torch.nn.Sequential(\n                    T.FrequencyMasking(freq_mask_param=N_MELS // 8),\n                    T.TimeMasking(time_mask_param=int(N_FRAMES * 0.1))\n                )\n            # self.spec_augmenter = nn.Sequential( # Alternative definition\n            #     T.FrequencyMasking(freq_mask_param=N_MELS // 8), # Mask up to 1/8th of mel bands\n            #     T.TimeMasking(time_mask_param=int(N_FRAMES * 0.1)) # Mask up to 10% of time steps\n            # )\n        else:\n            self.spec_augmenter = None\n\n        if self.preload_in_ram:\n            self.spectrograms = []\n            self.labels = []\n            print(f\"Preloading {len(dataframe)} audio files into RAM for {'training' if is_train else 'validation'}...\")\n            for idx in tqdm(range(len(dataframe))):\n                row = self.dataframe.iloc[idx]\n                file_path = row['full_path']\n                label = row['label_encoded']\n                # For preloading, we don't apply instance-specific augmentations like SpecAugment here,\n                # as they should be random for each epoch. Audio-level augs are applied during loading.\n                spectrogram = load_and_preprocess_audio(file_path, is_train=self.is_train) # is_train for audio augs\n                spectrogram_tensor = torch.tensor(spectrogram, dtype=torch.float32).unsqueeze(0)\n                spectrogram_tensor_3channel = spectrogram_tensor.repeat(3, 1, 1)\n                self.spectrograms.append(spectrogram_tensor_3channel)\n                self.labels.append(torch.tensor(label, dtype=torch.long))\n            print(f\"All {len(self.spectrograms)} spectrograms loaded into memory!\")\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        if self.preload_in_ram:\n            spectrogram_3channel = self.spectrograms[idx]\n            label = self.labels[idx]\n        else:\n            row = self.dataframe.iloc[idx]\n            file_path = row['full_path']\n            label_val = row['label_encoded']\n\n            # Load and process audio (includes audio-level augmentations if is_train)\n            spectrogram = load_and_preprocess_audio(file_path, is_train=self.is_train)\n            spectrogram_tensor = torch.tensor(spectrogram, dtype=torch.float32).unsqueeze(0) # (1, H, W)\n            spectrogram_3channel = spectrogram_tensor.repeat(3, 1, 1) # (3, H, W)\n            label = torch.tensor(label_val, dtype=torch.long)\n\n        # Apply SpecAugment (if training and enabled)\n        if self.is_train and self.spec_augmenter and random.random() < APPLY_AUGMENTATION_PROB:\n            # Ensure spec_augmenter is on the same device if it contains parameters,\n            # or apply before moving tensor to GPU in training loop.\n            # For torchaudio.transforms, they are typically stateless or handle device internally.\n            try:\n                spectrogram_3channel = self.spec_augmenter(spectrogram_3channel)\n            except Exception as e: # Catch potential errors with SpecAugment on edge cases\n                # print(f\"SpecAugment error for sample {idx}, shape {spectrogram_3channel.shape}: {e}\")\n                pass # Skip SpecAugment for this sample if it errors\n\n        return spectrogram_3channel, label\n\n#train_data = BirdSoundDataset(train_df, is_train=True, preload_in_ram=PRELOAD_DATA_IN_RAM)\n#val_data = BirdSoundDataset(val_df, is_train=False, preload_in_ram=PRELOAD_DATA_IN_RAM, apply_spec_augment=False)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:32:22.471392Z","iopub.execute_input":"2025-05-08T12:32:22.471768Z","iopub.status.idle":"2025-05-08T12:32:22.480644Z","shell.execute_reply.started":"2025-05-08T12:32:22.471745Z","shell.execute_reply":"2025-05-08T12:32:22.479947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = torch.load('/kaggle/input/train-databridclef/train_dataset.pt')\nval_data = torch.load(\"/kaggle/input/notebook863a8ddda3/val_dataset.pt\")\n\ntrain_loader = DataLoader(\n    train_data, batch_size=BATCH_SIZE, shuffle=True,\n    num_workers=NUM_WORKERS, pin_memory=True, drop_last=True # drop_last can be useful\n)\nval_loader = DataLoader(\n    val_data, batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=NUM_WORKERS, pin_memory=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:32:24.895103Z","iopub.execute_input":"2025-05-08T12:32:24.895719Z","iopub.status.idle":"2025-05-08T12:35:01.171319Z","shell.execute_reply.started":"2025-05-08T12:32:24.895695Z","shell.execute_reply":"2025-05-08T12:35:01.170571Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 6. Build Custom Model with EfficientNet Backbone ---\nprint(f\"\\nBuilding custom model with {MODEL_NAME} backbone...\")\n\n# Define a custom model class that uses EfficientNet as a backbone\nclass BirdClassifier(nn.Module):\n    def __init__(self, backbone_name, num_classes, pretrained=True):\n        super(BirdClassifier, self).__init__()\n        \n        # Get the EfficientNet backbone\n        if backbone_name == \"efficientnet_b3\":\n            weights = models.EfficientNet_B3_Weights.DEFAULT if pretrained else None\n            backbone = models.efficientnet_b3(weights=weights)\n            backbone_features = backbone.features\n            num_ftrs = 1536  # For EfficientNet B3\n        elif backbone_name == \"efficientnet_b2\":\n            weights = models.EfficientNet_B2_Weights.DEFAULT if pretrained else None\n            backbone = models.efficientnet_b2(weights=weights)\n            backbone_features = backbone.features\n            num_ftrs = 1408  # For EfficientNet B2\n        else:\n            raise ValueError(f\"Backbone model {backbone_name} not supported yet\")\n        \n        # Extract the feature extractor (everything except the classifier)\n        self.backbone = backbone_features\n        \n        # Global Average Pooling\n        self.gap = nn.AdaptiveAvgPool2d(1)\n        \n        # Custom classifier with an additional hidden layer\n        self.classifier = nn.Sequential(\n            nn.Dropout(0.2),\n            nn.Linear(num_ftrs, 512),  # Add a hidden layer\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n        \n        # Softmax activation for final layer (not included in training since CrossEntropyLoss has it)\n        self.softmax = nn.Softmax(dim=1)\n    \n    def forward(self, x):\n        # Extract features using the backbone\n        x = self.backbone(x)\n        \n        # Global Average Pooling\n        x = self.gap(x)\n        x = torch.flatten(x, 1)\n        \n        # Classification head\n        x = self.classifier(x)\n        \n        # Note: We don't apply softmax during training since CrossEntropyLoss includes it\n        # We only apply it during inference or when raw probabilities are needed\n        return x\n    \n    def predict_proba(self, x):\n        # For getting probabilities during inference\n        logits = self.forward(x)\n        return self.softmax(logits)\n\n# Initialize the model\nmodel = BirdClassifier(MODEL_NAME, NUM_CLASSES, pretrained=True)\nmodel = model.to(device)\n\n# Fine-tuning strategy\n# Option 1: Unfreeze all layers for end-to-end fine-tuning\nfor param in model.parameters():\n    param.requires_grad = True\nprint(\"Fine-tuning: Training all layers.\")\n\n# Option 2: Freeze backbone and train only classifier for the first few epochs\n# Comment this section if you want to train all layers from the beginning\n# for name, param in model.backbone.named_parameters():\n#     param.requires_grad = False\n# print(\"Fine-tuning: Only training the classifier initially.\")\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Total model parameters: {total_params:,}\")\nprint(f\"Trainable parameters: {trainable_params:,}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:38:59.833255Z","iopub.execute_input":"2025-05-08T12:38:59.833553Z","iopub.status.idle":"2025-05-08T12:39:00.154365Z","shell.execute_reply.started":"2025-05-08T12:38:59.833529Z","shell.execute_reply":"2025-05-08T12:39:00.153588Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 7. Define Loss Function and Optimizer ---\n\n# Label Smoothing Cross Entropy\nclass LabelSmoothingCrossEntropy(nn.Module):\n    def __init__(self, num_classes, smoothing=0.0, dim=-1):\n        super(LabelSmoothingCrossEntropy, self).__init__()\n        self.confidence = 1.0 - smoothing\n        self.smoothing = smoothing\n        self.num_classes = num_classes\n        self.dim = dim\n\n    def forward(self, pred, target):\n        pred = pred.log_softmax(dim=self.dim)\n        with torch.no_grad():\n            true_dist = torch.zeros_like(pred)\n            true_dist.fill_(self.smoothing / (self.num_classes - 1))\n            true_dist.scatter_(1, target.data.unsqueeze(1), self.confidence)\n        return torch.mean(torch.sum(-true_dist * pred, dim=self.dim))\n\nif LABEL_SMOOTHING > 0.0 and NUM_CLASSES > 1: # Label smoothing only makes sense for >1 classes\n    criterion = LabelSmoothingCrossEntropy(num_classes=NUM_CLASSES, smoothing=LABEL_SMOOTHING).to(device)\n    print(f\"Using Label Smoothing Cross Entropy with smoothing={LABEL_SMOOTHING}\")\nelse:\n    criterion = nn.CrossEntropyLoss().to(device)\n    print(\"Using standard CrossEntropyLoss.\")\n\n# Optimizer: AdamW\n# For differential learning rates:\nparam_groups = [\n    {'params': model.backbone.parameters(), 'lr': LEARNING_RATE / 10}, # Backbone\n    {'params': model.classifier.parameters(), 'lr': LEARNING_RATE}     # Classifier\n]\noptimizer = optim.AdamW(param_groups, weight_decay=WEIGHT_DECAY)\nprint(f\"Using AdamW optimizer with LR={LEARNING_RATE}, Weight Decay={WEIGHT_DECAY}\")\nprint(f\"Backbone LR={LEARNING_RATE / 10}, Classifier LR={LEARNING_RATE}\")\n\n# Learning Rate Scheduler\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-6)\nprint(f\"Using CosineAnnealingLR scheduler with T_max={EPOCHS} epochs.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:40:08.399819Z","iopub.execute_input":"2025-05-08T12:40:08.400418Z","iopub.status.idle":"2025-05-08T12:40:08.409992Z","shell.execute_reply.started":"2025-05-08T12:40:08.400395Z","shell.execute_reply":"2025-05-08T12:40:08.409069Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 8. Training Loop ---\nprint(f\"\\nStarting training for {EPOCHS} epochs on {device}...\")\n\nhistory = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': [], 'lr': []}\nbest_val_accuracy = 0.0 # Save based on best validation accuracy\nbest_model_wts = None\nearly_stopping_patience = 7 # More patience if using aggressive schedulers or complex data\nepochs_no_improve = 0\n\nfor epoch in range(EPOCHS):\n    epoch_start_time = time.time()\n\n    # --- Training Phase ---\n    model.train()\n    running_loss = 0.0\n    correct_predictions_train = 0\n    total_samples_train = 0\n\n    train_pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS} [Train]\", leave=False)\n    for inputs, labels in train_pbar:\n        inputs, labels = inputs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        # Optional: Gradient Clipping\n        # torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer.step()\n\n        running_loss += loss.item() * inputs.size(0)\n        _, predicted = torch.max(outputs.data, 1)\n        total_samples_train += labels.size(0)\n        correct_predictions_train += (predicted == labels).sum().item()\n        train_pbar.set_postfix(loss=loss.item(), acc=correct_predictions_train/total_samples_train if total_samples_train > 0 else 0.0)\n\n    epoch_train_loss = running_loss / total_samples_train if total_samples_train > 0 else 0.0\n    epoch_train_acc = correct_predictions_train / total_samples_train if total_samples_train > 0 else 0.0\n    history['train_loss'].append(epoch_train_loss)\n    history['train_acc'].append(epoch_train_acc)\n    history['lr'].append(optimizer.param_groups[0]['lr'])\n\n    # --- Validation Phase ---\n    model.eval()\n    running_loss_val = 0.0\n    correct_predictions_val = 0\n    total_samples_val = 0\n    val_pbar = tqdm(val_loader, desc=f\"Epoch {epoch+1}/{EPOCHS} [Val]\", leave=False)\n\n    with torch.no_grad():\n        for inputs, labels in val_pbar:\n            inputs, labels = inputs.to(device), labels.to(device)\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            running_loss_val += loss.item() * inputs.size(0)\n            _, predicted = torch.max(outputs.data, 1)\n            total_samples_val += labels.size(0)\n            correct_predictions_val += (predicted == labels).sum().item()\n            val_pbar.set_postfix(loss=loss.item(), acc=correct_predictions_val/total_samples_val if total_samples_val > 0 else 0.0)\n\n    epoch_val_loss = running_loss_val / total_samples_val if total_samples_val > 0 else 0.0\n    epoch_val_acc = correct_predictions_val / total_samples_val if total_samples_val > 0 else 0.0\n    history['val_loss'].append(epoch_val_loss)\n    history['val_acc'].append(epoch_val_acc)\n\n    epoch_duration = time.time() - epoch_start_time\n    print(f\"Epoch {epoch+1}/{EPOCHS} | Duration: {epoch_duration:.2f}s | LR: {optimizer.param_groups[0]['lr']:.2e}\")\n    print(f\"  Train Loss: {epoch_train_loss:.4f} | Train Acc: {epoch_train_acc:.4f}\")\n    print(f\"  Val Loss:   {epoch_val_loss:.4f} | Val Acc:   {epoch_val_acc:.4f}\")\n\n    # LR Scheduler Step (per epoch for CosineAnnealingLR with T_max=EPOCHS)\n    if scheduler and not isinstance(scheduler, optim.lr_scheduler.ReduceLROnPlateau):\n        scheduler.step()\n    elif scheduler and isinstance(scheduler, optim.lr_scheduler.ReduceLROnPlateau):\n        scheduler.step(epoch_val_loss)\n\n\n    # --- Save Best Model & Early Stopping ---\n    if epoch_val_acc > best_val_accuracy:\n        print(f\"  Validation accuracy improved ({best_val_accuracy:.4f} --> {epoch_val_acc:.4f}). Saving model...\")\n        best_val_accuracy = epoch_val_acc\n        best_model_wts = copy.deepcopy(model.state_dict())\n        torch.save(model.state_dict(), f'best_custom_{MODEL_NAME}_model.pth')\n        epochs_no_improve = 0\n    else:\n        epochs_no_improve += 1\n        print(f\"  Validation accuracy did not improve for {epochs_no_improve} epoch(s).\")\n\n    if epochs_no_improve >= early_stopping_patience:\n        print(f\"\\nEarly stopping triggered after {epoch + 1} epochs as validation accuracy did not improve for {early_stopping_patience} epochs.\")\n        break\n\nprint(\"\\nTraining finished.\")\nif best_model_wts:\n    print(f\"Best validation accuracy: {best_val_accuracy:.4f}\")\n    # model.load_state_dict(best_model_wts) # Load best model for further use if needed\nelse:\n    print(\"No improvement in validation accuracy was observed, or training stopped early.\")\n\n\n# --- Plot training history ---\nfig, axs = plt.subplots(1, 3, figsize=(20, 6))\n\naxs[0].plot(history['train_loss'], label='Train Loss')\naxs[0].plot(history['val_loss'], label='Validation Loss')\naxs[0].set_title('Model Loss')\naxs[0].set_xlabel('Epoch')\naxs[0].set_ylabel('Loss')\naxs[0].legend(); axs[0].grid(True)\n\naxs[1].plot(history['train_acc'], label='Train Accuracy')\naxs[1].plot(history['val_acc'], label='Validation Accuracy')\naxs[1].set_title('Model Accuracy')\naxs[1].set_xlabel('Epoch')\naxs[1].set_ylabel('Accuracy')\naxs[1].legend(); axs[1].grid(True)\n\naxs[2].plot(history['lr'], label='Learning Rate')\naxs[2].set_title('Learning Rate Schedule')\naxs[2].set_xlabel('Epoch')\naxs[2].set_ylabel('Learning Rate')\naxs[2].legend(); axs[2].grid(True)\n\nplt.tight_layout()\nplt.savefig(f\"training_history_custom_{MODEL_NAME}.png\")\nplt.show()\n\nprint(f\"Best model saved as: best_custom_{MODEL_NAME}_model.pth\")\nprint(f\"Training plots saved as: training_history_custom_{MODEL_NAME}.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:40:09.936759Z","iopub.execute_input":"2025-05-08T12:40:09.937045Z","iopub.status.idle":"2025-05-08T14:02:10.518334Z","shell.execute_reply.started":"2025-05-08T12:40:09.937022Z","shell.execute_reply":"2025-05-08T14:02:10.517377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}