{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# BirdCLEF 2025: AST Training Notebook - Modified for Full Training\n# =============================================================\n\n# Install necessary packages\n!pip install -q librosa==0.10.1 torchaudio timm soundfile audiomentations\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom timm.models.layers import to_2tuple, trunc_normal_\nimport matplotlib.pyplot as plt\nimport random\nimport warnings\nimport time\nfrom tqdm.auto import tqdm\nimport librosa\nimport cv2\nimport json\n\n# Silence warnings\nwarnings.filterwarnings(\"ignore\")\nprint(\"Torch version:\", torch.__version__)\nprint(\"Timm version:\", timm.__version__)\n\n# Set random seed for reproducibility\nSEED = 42\ndef set_seed(seed=SEED):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    \nset_seed()\n\n# Configuration\nclass CFG:\n    # Paths\n    train_audio_dir = '/kaggle/input/birdclef-2025/train_audio'\n    train_csv = '/kaggle/input/birdclef-2025/train.csv'\n    taxonomy_csv = '/kaggle/input/birdclef-2025/taxonomy.csv'\n    output_dir = '/kaggle/working/ast_model'\n    \n    # Audio parameters\n    sample_rate = 32000\n    duration = 5  # seconds per clip\n    \n    # Mel spectrogram parameters\n    n_mels = 128\n    n_fft = 1024\n    hop_length = 512\n    fmin = 50\n    fmax = 14000\n    \n    # AST model parameters\n    fstride = 10\n    tstride = 10\n    patch_size = 16\n    model_size = 'base224'  # Changed from base384 to base224 for better compatibility\n    \n    # Vision Transformer expected input size\n    target_height = 224\n    target_width = 224\n    \n    # Training parameters\n    batch_size = 16\n    epochs = 10\n    lr = 1e-4\n    weight_decay = 1e-6\n    \n    # Augmentation parameters\n    freqm = 24  # Frequency mask max length\n    timem = 96  # Time mask max length\n    mixup_alpha = 0.5\n    \n    # Device\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    # Debug mode - set to True to train on a small subset of data\n    debug = False\n\nif CFG.debug:\n    CFG.epochs = 2\n    CFG.batch_size = 8\n\nos.makedirs(CFG.output_dir, exist_ok=True)\nprint(f\"Using device: {CFG.device}\")\n\n# Load data\nprint(\"Loading data...\")\ndf = pd.read_csv(CFG.train_csv)\ntaxonomy_df = pd.read_csv(CFG.taxonomy_csv)\n\n# Add filepath column\ndf['filepath'] = df['filename'].apply(lambda x: os.path.join(CFG.train_audio_dir, x))\n\n# Create a dictionary mapping from primary_label to index\n# Important: We'll use it consistently for both training and inference\nunique_labels = taxonomy_df['primary_label'].unique()\nlabel_map = {label: idx for idx, label in enumerate(unique_labels)}\nnum_classes = len(label_map)\nprint(f\"Number of classes: {num_classes}\")\n\n# Add label_idx column - make sure to use our consistent label mapping\ndf['label_idx'] = df['primary_label'].map(label_map)\n\n# Check for any NaN values in label_idx column\nif df['label_idx'].isna().any():\n    print(f\"WARNING: Found {df['label_idx'].isna().sum()} rows with NaN label_idx\")\n    print(\"Example rows with missing labels:\")\n    print(df[df['label_idx'].isna()].head())\n    \n    # Remove rows with NaN label_idx\n    df = df.dropna(subset=['label_idx']).reset_index(drop=True)\n    print(f\"Removed rows with missing labels. New shape: {df.shape}\")\n\n# Convert label_idx to integer\ndf['label_idx'] = df['label_idx'].astype(int)\n\n# For debug mode, use smaller dataset\nif CFG.debug:\n    df = df.sample(min(500, len(df)), random_state=SEED).reset_index(drop=True)\n\nprint(f\"Training data shape: {df.shape}\")\nprint(df.head())\n\n# Define AST model\nclass PatchEmbed(nn.Module):\n    \"\"\"2D Image to Patch Embedding\"\"\"\n    def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):\n        super().__init__()\n        \n        img_size = to_2tuple(img_size)\n        patch_size = to_2tuple(patch_size)\n        num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0])\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.num_patches = num_patches\n        \n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n        \n    def forward(self, x):\n        x = self.proj(x).flatten(2).transpose(1, 2)\n        return x\n\nclass ASTModel(nn.Module):\n    \"\"\"Audio Spectrogram Transformer model\"\"\"\n    def __init__(self, label_dim=527, fstride=10, tstride=10, input_fdim=128, input_tdim=1024, \n                 imagenet_pretrain=True, model_size='base224'):\n        super(ASTModel, self).__init__()\n        \n        # Override timm input shape restriction\n        timm.models.vision_transformer.PatchEmbed = PatchEmbed\n        \n        # Print model configuration\n        print(f'AST Model: size={model_size}, input_fdim={input_fdim}, input_tdim={input_tdim}')\n        print(f'frequency stride={fstride}, time stride={tstride}')\n        \n        # Load model - first check what models are available\n        available_models = timm.list_models('*vit*')\n        print(f\"Available models containing 'vit': {available_models[:5]} and {len(available_models)} more...\")\n        \n        # Try to use a model that's available\n        if model_size == 'tiny224':\n            try:\n                self.v = timm.create_model('vit_deit_tiny_distilled_patch16_224', pretrained=imagenet_pretrain)\n            except RuntimeError:\n                # Fallback to a model that should be available\n                print(\"Falling back to vit_tiny_patch16_224\")\n                self.v = timm.create_model('vit_tiny_patch16_224', pretrained=imagenet_pretrain)\n        elif model_size == 'small224':\n            try:\n                self.v = timm.create_model('vit_deit_small_distilled_patch16_224', pretrained=imagenet_pretrain)\n            except RuntimeError:\n                print(\"Falling back to vit_small_patch16_224\")\n                self.v = timm.create_model('vit_small_patch16_224', pretrained=imagenet_pretrain)\n        elif model_size == 'base224':\n            try:\n                self.v = timm.create_model('vit_deit_base_distilled_patch16_224', pretrained=imagenet_pretrain)\n            except RuntimeError:\n                print(\"Falling back to vit_base_patch16_224\")\n                self.v = timm.create_model('vit_base_patch16_224', pretrained=imagenet_pretrain)\n        elif model_size == 'base384':\n            try:\n                self.v = timm.create_model('vit_deit_base_distilled_patch16_384', pretrained=imagenet_pretrain)\n            except RuntimeError:\n                print(\"Falling back to vit_base_patch16_384\")\n                try:\n                    self.v = timm.create_model('vit_base_patch16_384', pretrained=imagenet_pretrain)\n                except RuntimeError:\n                    print(\"Falling back to vit_base_patch16_224\")\n                    self.v = timm.create_model('vit_base_patch16_224', pretrained=imagenet_pretrain)\n        else:\n            raise Exception('Model size must be one of tiny224, small224, base224, base384.')\n            \n        # Check if model has distillation token\n        self.has_dist_token = hasattr(self.v, 'dist_token')\n        print(f\"Model has distillation token: {self.has_dist_token}\")\n        \n        self.original_num_patches = self.v.patch_embed.num_patches\n        self.oringal_hw = int(self.original_num_patches ** 0.5)\n        self.original_embedding_dim = self.v.pos_embed.shape[2]\n        self.mlp_head = nn.Sequential(nn.LayerNorm(self.original_embedding_dim), \n                                     nn.Linear(self.original_embedding_dim, label_dim))\n        \n        # Get shape automatically\n        f_dim, t_dim = self.get_shape(fstride, tstride, input_fdim, input_tdim)\n        num_patches = f_dim * t_dim\n        self.v.patch_embed.num_patches = num_patches\n        \n        print(f'number of patches={num_patches}')\n            \n        # Linear projection\n        new_proj = torch.nn.Conv2d(1, self.original_embedding_dim, kernel_size=(16, 16), stride=(fstride, tstride))\n        if imagenet_pretrain:\n            new_proj.weight = torch.nn.Parameter(torch.sum(self.v.patch_embed.proj.weight, dim=1).unsqueeze(1))\n            new_proj.bias = self.v.patch_embed.proj.bias\n        self.v.patch_embed.proj = new_proj\n        \n        # Positional embedding\n        if imagenet_pretrain:\n            # Get the positional embedding from model\n            if self.has_dist_token:\n                new_pos_embed = self.v.pos_embed[:, 2:, :].detach().reshape(1, self.original_num_patches, self.original_embedding_dim).transpose(1, 2).reshape(1, self.original_embedding_dim, self.oringal_hw, self.oringal_hw)\n            else:\n                new_pos_embed = self.v.pos_embed[:, 1:, :].detach().reshape(1, self.original_num_patches, self.original_embedding_dim).transpose(1, 2).reshape(1, self.original_embedding_dim, self.oringal_hw, self.oringal_hw)\n            \n            # Cut or interpolate position embedding\n            if t_dim <= self.oringal_hw:\n                new_pos_embed = new_pos_embed[:, :, :, int(self.oringal_hw / 2) - int(t_dim / 2): int(self.oringal_hw / 2) - int(t_dim / 2) + t_dim]\n            else:\n                new_pos_embed = torch.nn.functional.interpolate(new_pos_embed, size=(self.oringal_hw, t_dim), mode='bilinear')\n                \n            # Cut or interpolate position embedding\n            if f_dim <= self.oringal_hw:\n                new_pos_embed = new_pos_embed[:, :, int(self.oringal_hw / 2) - int(f_dim / 2): int(self.oringal_hw / 2) - int(f_dim / 2) + f_dim, :]\n            else:\n                new_pos_embed = torch.nn.functional.interpolate(new_pos_embed, size=(f_dim, t_dim), mode='bilinear')\n                \n            # Flatten the position embedding\n            new_pos_embed = new_pos_embed.reshape(1, self.original_embedding_dim, num_patches).transpose(1, 2)\n            \n            # Concatenate with cls token and distillation token\n            if self.has_dist_token:\n                self.v.pos_embed = nn.Parameter(torch.cat([self.v.pos_embed[:, :2, :].detach(), new_pos_embed], dim=1))\n            else:\n                self.v.pos_embed = nn.Parameter(torch.cat([self.v.pos_embed[:, :1, :].detach(), new_pos_embed], dim=1))\n        else:\n            # Random initialization\n            if self.has_dist_token:\n                new_pos_embed = nn.Parameter(torch.zeros(1, self.v.patch_embed.num_patches + 2, self.original_embedding_dim))\n            else:\n                new_pos_embed = nn.Parameter(torch.zeros(1, self.v.patch_embed.num_patches + 1, self.original_embedding_dim))\n            self.v.pos_embed = new_pos_embed\n            trunc_normal_(self.v.pos_embed, std=.02)\n        \n    def get_shape(self, fstride, tstride, input_fdim=128, input_tdim=1024):\n        test_input = torch.randn(1, 1, input_fdim, input_tdim)\n        test_proj = nn.Conv2d(1, self.original_embedding_dim, kernel_size=(16, 16), stride=(fstride, tstride))\n        test_out = test_proj(test_input)\n        f_dim = test_out.shape[2]\n        t_dim = test_out.shape[3]\n        return f_dim, t_dim\n    \n    def forward(self, x):\n        \"\"\"\n        :param x: Input spectrogram, expected shape: (batch_size, time_frame_num, frequency_bins)\n        :return: prediction\n        \"\"\"\n        # Input shape: (batch_size, time_frame_num, frequency_bins)\n        x = x.unsqueeze(1)        # Add channel dimension: (B, 1, T, F)\n        x = x.transpose(2, 3)     # -> (B, 1, F, T)\n        \n        B = x.shape[0]\n        x = self.v.patch_embed(x)\n        \n        # Handle both model types (with and without distillation token)\n        if self.has_dist_token:\n            cls_tokens = self.v.cls_token.expand(B, -1, -1)\n            dist_token = self.v.dist_token.expand(B, -1, -1)\n            x = torch.cat((cls_tokens, dist_token, x), dim=1)\n        else:\n            cls_tokens = self.v.cls_token.expand(B, -1, -1)\n            x = torch.cat((cls_tokens, x), dim=1)\n            \n        x = x + self.v.pos_embed\n        x = self.v.pos_drop(x)\n        \n        for blk in self.v.blocks:\n            x = blk(x)\n            \n        x = self.v.norm(x)\n        \n        # Handle both model types for output\n        if self.has_dist_token:\n            x = (x[:, 0] + x[:, 1]) / 2  # Average of cls and dist token\n        else:\n            x = x[:, 0]  # Just use cls token\n        \n        x = self.mlp_head(x)\n        return x\n\n# Audio processing functions\ndef audio_to_melspec(audio_data, cfg):\n    \"\"\"Convert audio data to mel spectrogram\"\"\"\n    # Handle NaN values\n    if np.isnan(audio_data).any():\n        mean_signal = np.nanmean(audio_data)\n        audio_data = np.nan_to_num(audio_data, nan=mean_signal)\n    \n    # Generate mel spectrogram\n    mel_spec = librosa.feature.melspectrogram(\n        y=audio_data,\n        sr=cfg.sample_rate,\n        n_fft=cfg.n_fft,\n        hop_length=cfg.hop_length,\n        n_mels=cfg.n_mels,\n        fmin=cfg.fmin,\n        fmax=cfg.fmax,\n        power=2.0\n    )\n    \n    # Convert to dB scale\n    mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max)\n    \n    # Normalize to [0, 1]\n    mel_spec_norm = (mel_spec_db - mel_spec_db.min()) / (mel_spec_db.max() - mel_spec_db.min() + 1e-8)\n    \n    return mel_spec_norm\n\ndef process_audio_file(audio_path, cfg):\n    \"\"\"Process a single audio file to get the mel spectrogram\"\"\"\n    try:\n        # Load audio\n        audio_data, _ = librosa.load(audio_path, sr=cfg.sample_rate)\n        \n        # Calculate target length in samples\n        target_length = int(cfg.duration * cfg.sample_rate)\n        \n        # Handle audio shorter than target duration\n        if len(audio_data) < target_length:\n            # Pad with zeros\n            audio_data = np.pad(audio_data, \n                             (0, target_length - len(audio_data)),\n                             mode='constant')\n        \n        # Take center segment if longer than target duration\n        if len(audio_data) > target_length:\n            start_idx = (len(audio_data) - target_length) // 2\n            audio_data = audio_data[start_idx:start_idx + target_length]\n        \n        # Generate mel spectrogram\n        mel_spec = audio_to_melspec(audio_data, cfg)\n        \n        return mel_spec.astype(np.float32)\n    \n    except Exception as e:\n        print(f\"Error processing {audio_path}: {e}\")\n        return None\n\n# Resize spectrogram to match model's expected input dimensions\ndef resize_spectrogram(spec, target_height=224, target_width=224):\n    \"\"\"Resize a spectrogram to match the expected input dimensions of the model\"\"\"\n    # Convert to numpy if it's a tensor\n    if isinstance(spec, torch.Tensor):\n        spec_np = spec.numpy()\n    else:\n        spec_np = spec\n    \n    # Ensure the input is 2D\n    if spec_np.ndim != 2:\n        print(f\"Warning: Expected 2D spectrogram, got shape {spec_np.shape}\")\n        if spec_np.ndim == 3 and spec_np.shape[0] == 1:\n            spec_np = spec_np.squeeze(0)  # Remove singleton dimension\n    \n    # Resize using OpenCV\n    resized_spec = cv2.resize(spec_np, (target_width, target_height))\n    \n    # Convert back to tensor if input was a tensor\n    if isinstance(spec, torch.Tensor):\n        return torch.tensor(resized_spec, dtype=torch.float32)\n    else:\n        return resized_spec.astype(np.float32)\n\n# Dataset class\nclass BirdCLEFDataset(Dataset):\n    def __init__(self, df, cfg, transform=None):\n        self.df = df\n        self.cfg = cfg\n        self.transform = transform\n        self.num_classes = num_classes  # Use the global num_classes value\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        # Get audio file path and label\n        audio_path = row['filepath']\n        label_idx = row['label_idx']\n        \n        # Process audio file\n        spec = process_audio_file(audio_path, self.cfg)\n        \n        if spec is None:\n            # Return zeros if processing failed\n            spec = np.zeros((self.cfg.n_mels, int(self.cfg.duration * self.cfg.sample_rate / self.cfg.hop_length) + 1), dtype=np.float32)\n        \n        # Resize spectrogram to match model's expected input dimensions\n        spec = resize_spectrogram(spec, target_height=self.cfg.target_height, target_width=self.cfg.target_width)\n        \n        # Convert to tensor\n        spec_tensor = torch.tensor(spec, dtype=torch.float32)\n        \n        # Apply transformations if any\n        if self.transform:\n            spec_tensor = self.transform(spec_tensor)\n            \n        # Create one-hot encoded label tensor\n        label_tensor = torch.zeros(self.num_classes, dtype=torch.float32)\n        label_tensor[label_idx] = 1.0\n        \n        # Add secondary labels if they exist\n        if 'secondary_labels' in row and row['secondary_labels'] not in ['[]', '', None]:\n            if isinstance(row['secondary_labels'], str):\n                try:\n                    secondary_labels = eval(row['secondary_labels'])\n                    for label in secondary_labels:\n                        if label in label_map:\n                            sec_idx = label_map[label]\n                            if 0 <= sec_idx < self.num_classes:  # Ensure the index is valid\n                                label_tensor[sec_idx] = 1.0\n                except:\n                    pass  # Skip if there's an error parsing secondary labels\n        \n        return {\n            'spectrogram': spec_tensor,\n            'label': label_tensor\n        }\n\n# Augmentation transforms\nclass SpecAugment:\n    def __init__(self, freq_mask=24, time_mask=96):\n        self.freq_mask = freq_mask\n        self.time_mask = time_mask\n        \n    def __call__(self, spec):\n        # Apply frequency masking\n        if self.freq_mask > 0:\n            num_masks = random.randint(1, 2)\n            for _ in range(num_masks):\n                freq_mask_size = random.randint(1, self.freq_mask)\n                freq_start = random.randint(0, spec.shape[0] - freq_mask_size)\n                spec[freq_start:freq_start + freq_mask_size, :] = 0\n        \n        # Apply time masking\n        if self.time_mask > 0:\n            num_masks = random.randint(1, 2)\n            for _ in range(num_masks):\n                time_mask_size = random.randint(1, self.time_mask)\n                time_start = random.randint(0, spec.shape[1] - time_mask_size)\n                spec[:, time_start:time_start + time_mask_size] = 0\n                \n        # Apply random brightness\n        if random.random() < 0.5:\n            gain = random.uniform(0.8, 1.2)\n            bias = random.uniform(-0.1, 0.1)\n            spec = torch.clamp(gain * spec + bias, 0, 1)\n            \n        return spec\n\n# Mixup function\ndef mixup_data(x, y, alpha=0.4):\n    \"\"\"Applies mixup augmentation to the batch\"\"\"\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1\n\n    batch_size = x.size(0)\n    index = torch.randperm(batch_size).to(x.device)\n\n    mixed_x = lam * x + (1 - lam) * x[index]\n    mixed_y = lam * y + (1 - lam) * y[index]\n    \n    return mixed_x, mixed_y\n\n# No validation split - use all data for training\nprint(f\"Training on all data: {len(df)} samples\")\n\n# Create dataset and dataloader\ntrain_transform = SpecAugment(freq_mask=CFG.freqm, time_mask=CFG.timem)\ntrain_dataset = BirdCLEFDataset(df, CFG, transform=train_transform)\n\n# Check the shape of a sample item to ensure it's properly resized\nsample_item = train_dataset[0]\nprint(f\"Sample spectrogram shape: {sample_item['spectrogram'].shape}\")\nprint(f\"Sample label shape: {sample_item['label'].shape}\")\n\ntrain_loader = DataLoader(\n    train_dataset, \n    batch_size=CFG.batch_size, \n    shuffle=True, \n    num_workers=2, \n    pin_memory=True\n)\n\nprint(f\"Train dataloader: {len(train_loader)} batches\")\n\n# Initialize model\nmodel = ASTModel(\n    label_dim=num_classes,\n    fstride=CFG.fstride,\n    tstride=CFG.tstride,\n    input_fdim=CFG.target_height,  # Use target dimensions that match what we're resizing to\n    input_tdim=CFG.target_width,\n    imagenet_pretrain=True,\n    model_size=CFG.model_size\n)\n\nmodel = model.to(CFG.device)\n\n# Loss and optimizer\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = optim.Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=CFG.lr/100)\n\n# Training function\ndef train_one_epoch(model, loader, criterion, optimizer, device, mixup_alpha=0.5):\n    model.train()\n    losses = []\n    pbar = tqdm(loader, desc='Training')\n    \n    for batch in pbar:\n        # Get data\n        specs = batch['spectrogram'].to(device)\n        labels = batch['label'].to(device)\n        \n        # Apply mixup with probability 0.5\n        if mixup_alpha > 0 and random.random() < 0.5:\n            specs, labels = mixup_data(specs, labels, alpha=mixup_alpha)\n        \n        # Forward pass\n        optimizer.zero_grad()\n        logits = model(specs)\n        loss = criterion(logits, labels)\n        \n        # Backward pass\n        loss.backward()\n        optimizer.step()\n        \n        # Record loss\n        losses.append(loss.item())\n        \n        # Update progress bar\n        pbar.set_postfix({'loss': np.mean(losses[-10:])})\n    \n    return np.mean(losses)\n\n# Training loop\nhistory = {'train_loss': []}\n\n# Save initial model configuration\nmodel_config = {}\nfor k, v in CFG.__dict__.items():\n    if not k.startswith('__') and not callable(v):\n        # Convert non-serializable objects to strings\n        if k == 'device':\n            model_config[k] = str(v)\n        else:\n            model_config[k] = v\n\nwith open(os.path.join(CFG.output_dir, 'model_config.json'), 'w') as f:\n    json.dump(model_config, f)\n\nfor epoch in range(1, CFG.epochs + 1):\n    print(f\"\\nEpoch {epoch}/{CFG.epochs}\")\n    \n    # Train\n    train_loss = train_one_epoch(\n        model=model,\n        loader=train_loader,\n        criterion=criterion,\n        optimizer=optimizer,\n        device=CFG.device,\n        mixup_alpha=CFG.mixup_alpha\n    )\n    \n    # Update learning rate\n    scheduler.step()\n    \n    # Update history\n    history['train_loss'].append(train_loss)\n    \n    # Print metrics\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"LR: {optimizer.param_groups[0]['lr']:.6f}\")\n    \n    # Save checkpoint every epoch\n    checkpoint_path = os.path.join(CFG.output_dir, f'model_epoch_{epoch}.pth')\n    \n    # Create a serializable config dictionary\n    config_dict = {}\n    for k, v in CFG.__dict__.items():\n        if not k.startswith('__') and not callable(v):\n            # Convert non-serializable objects to strings\n            if k == 'device':\n                config_dict[k] = str(v)\n            else:\n                config_dict[k] = v\n                \n    torch.save({\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'scheduler_state_dict': scheduler.state_dict(),\n        'epoch': epoch,\n        'num_classes': num_classes,\n        'label_map': label_map,\n        'config': config_dict\n    }, checkpoint_path)\n    \n    print(f\"Saved checkpoint to {checkpoint_path}\")\n\n# Save final model\nconfig_dict = {}\nfor k, v in CFG.__dict__.items():\n    if not k.startswith('__') and not callable(v):\n        # Convert non-serializable objects to strings\n        if k == 'device':\n            config_dict[k] = str(v)\n        else:\n            config_dict[k] = v\n            \ntorch.save({\n    'model_state_dict': model.state_dict(),\n    'optimizer_state_dict': optimizer.state_dict(),\n    'scheduler_state_dict': scheduler.state_dict(),\n    'epoch': CFG.epochs,\n    'num_classes': num_classes,\n    'label_map': label_map,\n    'config': config_dict\n}, os.path.join(CFG.output_dir, 'final_model.pth'))\n\n# Save training history\nwith open(os.path.join(CFG.output_dir, 'history.json'), 'w') as f:\n    json.dump(history, f)\n    \nprint(f\"Training complete!\")\n\n# Plot training history\nplt.figure(figsize=(8, 4))\nplt.plot(history['train_loss'], label='Train Loss')\nplt.title('Training Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.savefig(os.path.join(CFG.output_dir, 'training_history.png'))\nplt.show()\n\nprint(f\"Model saved to {CFG.output_dir}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-14T20:51:12.271306Z","iopub.execute_input":"2025-05-14T20:51:12.271605Z","iopub.status.idle":"2025-05-15T04:10:15.037964Z","shell.execute_reply.started":"2025-05-14T20:51:12.271581Z","shell.execute_reply":"2025-05-15T04:10:15.036912Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}