{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        # Check if the file extension is not '.jpg'\n        if not filename.lower().endswith('.jpg'):\n            print(os.path.join(dirname, filename))","metadata":{"_uuid":"b06fc474-39bb-41ad-9f5e-ea77950d0b2f","_cell_guid":"6febeefe-8926-4d35-8537-3a70ca6741a8","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2024-11-28T19:47:34.468990Z","iopub.execute_input":"2024-11-28T19:47:34.469910Z","iopub.status.idle":"2024-11-28T19:47:40.103555Z","shell.execute_reply.started":"2024-11-28T19:47:34.469846Z","shell.execute_reply":"2024-11-28T19:47:40.102864Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\nimport random\nfrom pathlib import Path\nfrom typing import Optional, Tuple, Dict\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nimport cv2\nfrom sklearn.model_selection import StratifiedKFold\nfrom torch.cuda.amp import autocast, GradScaler\nfrom tqdm import tqdm\n\nfrom albumentations import (\n    Compose, RandomResizedCrop, Transpose, HorizontalFlip, VerticalFlip,\n    ShiftScaleRotate, HueSaturationValue, RandomBrightnessContrast,\n    CoarseDropout, Normalize, CenterCrop, Resize\n)\nfrom albumentations.pytorch import ToTensorV2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:40.105093Z","iopub.execute_input":"2024-11-28T19:47:40.105381Z","iopub.status.idle":"2024-11-28T19:47:44.084458Z","shell.execute_reply.started":"2024-11-28T19:47:40.105349Z","shell.execute_reply":"2024-11-28T19:47:44.083586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------------------------------\n# Define paths for data and outputs\n# ---------------------------------------------------\nPATHS = {\n    'TRAIN_CSV': '/kaggle/input/cassava-leaf-disease-classification/train.csv',  # Path to training CSV\n    'TEST_CSV': '/kaggle/input/cassava-leaf-disease-classification/sample_submission.csv',  # Path to test submission CSV\n    'DISEASE_MAP': '/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json',  # Path to disease mapping JSON\n    'TRAIN_IMAGES': '/kaggle/input/cassava-leaf-disease-classification/train_images',  # Directory containing training images\n    'TEST_IMAGES': '/kaggle/input/cassava-leaf-disease-classification/test_images',  # Directory containing test images\n    'OUTPUT': '/kaggle/working/submission.csv',  # Path to save the final submission\n    'WEIGHTS': '/kaggle/working/weights'  # Directory to save model weights\n}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.085449Z","iopub.execute_input":"2024-11-28T19:47:44.085897Z","iopub.status.idle":"2024-11-28T19:47:44.090554Z","shell.execute_reply.started":"2024-11-28T19:47:44.085867Z","shell.execute_reply":"2024-11-28T19:47:44.089629Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    \"\"\"Configuration class for ViT model training and inference.\"\"\"\n    \n    def __init__(self):\n        # Check CUDA availability\n        self.device_type = 'cuda' if torch.cuda.is_available() else 'cpu'  # Device type\n        self.device = torch.device(self.device_type)  # PyTorch device\n        \n        # Model Configuration\n        self.model_name: str = 'vit_base_patch16_384'  # Vision Transformer model name\n        self.image_size: int = 384  # Input image size\n        self.patch_size: int = 16  # Patch size for ViT\n        self.hidden_size: int = 768  # Hidden size of the transformer\n        self.num_heads: int = 12  # Number of attention heads\n        self.num_layers: int = 12  # Number of transformer layers\n        self.pretrained: bool = True  # Whether to use pretrained weights\n        \n        # Training Configuration\n        self.seed: int = 719  # Random seed for reproducibility\n        self.num_epochs: int = 10  # Number of training epochs\n        self.train_batch_size: int = 8 if self.device_type == 'cuda' else 4  # Training batch size based on device\n        self.valid_batch_size: int = 16 if self.device_type == 'cuda' else 8  # Validation batch size based on device\n        self.learning_rate: float = 1e-4  # Learning rate for optimizer\n        self.weight_decay: float = 0.01  # Weight decay for optimizer\n        self.num_workers: int = 4 if self.device_type == 'cuda' else 2  # Number of workers for data loading\n        self.grad_accum_steps: int = 2  # Gradient accumulation steps\n        \n        # Mixed Precision\n        self.fp16: bool = self.device_type == 'cuda'  # Use mixed precision only if CUDA is available\n        \n        # Cross Validation\n        self.num_folds: int = 5  # Number of cross-validation folds\n        self.tta_steps: int = 3  # Number of Test Time Augmentation steps\n        self.used_epochs: list = [7, 8, 9]  # Epochs to use for inference\n        self.used_folds: list = [0, 2, 3]  # Folds to use for inference\n        \n        # Normalization\n        self.mean: list = [0.485, 0.456, 0.406]  # Mean values for normalization (ImageNet)\n        self.std: list = [0.229, 0.224, 0.225]  # Standard deviation values for normalization (ImageNet)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.092266Z","iopub.execute_input":"2024-11-28T19:47:44.092530Z","iopub.status.idle":"2024-11-28T19:47:44.105898Z","shell.execute_reply.started":"2024-11-28T19:47:44.092504Z","shell.execute_reply":"2024-11-28T19:47:44.105157Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CassavaDataset(Dataset):\n    \"\"\"Dataset class for Cassava Leaf Disease Classification.\"\"\"\n    \n    def __init__(\n        self,\n        df: pd.DataFrame,\n        data_root: str,\n        transforms: Optional[Compose] = None,\n        output_label: bool = True\n    ):\n        super().__init__()\n        self.df = df.reset_index(drop=True)  # Reset DataFrame index for consistency\n        self.transforms = transforms         # Data augmentation and preprocessing transforms\n        self.data_root = Path(data_root)     # Root directory for image data\n        self.output_label = output_label     # Flag to determine if labels are returned\n    \n    def __len__(self) -> int:\n        return len(self.df)  # Return the total number of samples\n    \n    def __getitem__(self, index: int) -> Tuple[torch.Tensor, Optional[int]]:\n        # Get image ID from the DataFrame\n        image_id = self.df.iloc[index]['image_id']\n        \n        # Ensure image has extension\n        if not image_id.lower().endswith(('.jpg', '.jpeg', '.png')):\n            image_id = f\"{image_id}.jpg\"  # Add .jpg extension if missing\n        \n        # Construct image path\n        image_path = self.data_root / image_id\n        \n        try:\n            image = self._load_image(str(image_path))  # Attempt to load the image\n        except Exception as e:\n            print(f\"Error loading image {image_path}: {str(e)}\")  # Print error message\n            # Return a blank image in case of error\n            image = np.zeros((384, 384, 3), dtype=np.uint8)\n        \n        if self.transforms:\n            image = self.transforms(image=image)['image']  # Apply transformations if any\n        \n        if self.output_label:\n            target = self.df.iloc[index]['label']  # Get label if output_label is True\n            return image, target  # Return image and label\n        return image  # Return only image for inference\n    \n    @staticmethod\n    def _load_image(path: str) -> np.ndarray:\n        \"\"\"Load and convert BGR image to RGB.\"\"\"\n        image = cv2.imread(path)  # Read image using OpenCV\n        if image is None:\n            raise ValueError(f\"Failed to load image at {path}\")  # Raise error if image is not found\n        return cv2.cvtColor(image, cv2.COLOR_BGR2RGB)            # Convert BGR to RGB format\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.107034Z","iopub.execute_input":"2024-11-28T19:47:44.107341Z","iopub.status.idle":"2024-11-28T19:47:44.120266Z","shell.execute_reply.started":"2024-11-28T19:47:44.107307Z","shell.execute_reply":"2024-11-28T19:47:44.119431Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def print_dataset_info(train_df: pd.DataFrame) -> None:\n    \"\"\"Print dataset information.\"\"\"\n    print(f\"Total training samples: {len(train_df)}\")      # Print total number of training samples\n    print(\"\\nLabel distribution:\")                         # Header for label distribution\n    print(train_df['label'].value_counts(normalize=True))  # Print normalized label counts\n    print(\"\\nSample image IDs:\")                           # Header for sample image IDs\n    print(train_df['image_id'].head())                     # Print first few image IDs\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.121295Z","iopub.execute_input":"2024-11-28T19:47:44.121533Z","iopub.status.idle":"2024-11-28T19:47:44.138927Z","shell.execute_reply.started":"2024-11-28T19:47:44.121510Z","shell.execute_reply":"2024-11-28T19:47:44.138125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed: int) -> None:\n    \"\"\"Set random seeds for reproducibility.\"\"\"\n    random.seed(seed)                          # Set Python random seed\n    os.environ['PYTHONHASHSEED'] = str(seed)   # Set environment variable for Python hash seed\n    np.random.seed(seed)                       # Set NumPy random seed\n    torch.manual_seed(seed)                    # Set PyTorch random seed\n    torch.cuda.manual_seed(seed)               # Set CUDA random seed\n    torch.backends.cudnn.deterministic = True  # Ensure deterministic behavior\n    torch.backends.cudnn.benchmark = True      # Enable benchmarking for performance\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.139656Z","iopub.execute_input":"2024-11-28T19:47:44.139905Z","iopub.status.idle":"2024-11-28T19:47:44.150520Z","shell.execute_reply.started":"2024-11-28T19:47:44.139881Z","shell.execute_reply":"2024-11-28T19:47:44.149742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CassavaViT(nn.Module):\n    \"\"\"Vision Transformer model for Cassava disease classification.\"\"\"\n    \n    def __init__(\n        self,\n        num_classes: int,\n        pretrained: bool = True,\n        model_name: str = 'vit_base_patch16_384'\n    ):\n        super().__init__()\n        self.model = timm.create_model(\n            model_name,\n            pretrained=pretrained,   # Use pretrained weights if True\n            num_classes=num_classes  # Set number of output classes\n        )\n    \n    def forward(self, x):\n        return self.model(x)  # Forward pass through the ViT model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.151440Z","iopub.execute_input":"2024-11-28T19:47:44.151679Z","iopub.status.idle":"2024-11-28T19:47:44.160086Z","shell.execute_reply.started":"2024-11-28T19:47:44.151656Z","shell.execute_reply":"2024-11-28T19:47:44.159455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ViTDataTransforms:\n    \"\"\"Data augmentation and preprocessing transforms optimized for ViT.\"\"\"\n    \n    @staticmethod\n    def get_train_transforms(config: Config) -> Compose:\n        \"\"\"Return training data augmentation transforms.\"\"\"\n        return Compose([\n            RandomResizedCrop(\n                height=config.image_size,\n                width=config.image_size,\n                scale=(0.8, 1.0)\n            ),  # Randomly crop and resize image\n            Transpose(p=0.5),       # Randomly transpose image dimensions\n            HorizontalFlip(p=0.5),  # Random horizontal flip\n            VerticalFlip(p=0.5),    # Random vertical flip\n            ShiftScaleRotate(\n                shift_limit=0.2,\n                scale_limit=0.2,\n                rotate_limit=30,\n                p=0.5\n            ),  # Random shift, scale, and rotation\n            HueSaturationValue(\n                hue_shift_limit=20,\n                sat_shift_limit=30,\n                val_shift_limit=20,\n                p=0.5\n            ),  # Randomly change hue, saturation, and value\n            RandomBrightnessContrast(\n                brightness_limit=0.2,\n                contrast_limit=0.2,\n                p=0.5\n            ),  # Random brightness and contrast adjustments\n            Normalize(\n                mean=config.mean,\n                std=config.std,\n                max_pixel_value=255.0,\n                p=1.0\n            ),  # Normalize image with mean and std\n            CoarseDropout(\n                max_holes=8,\n                max_height=config.image_size // 16,\n                max_width=config.image_size // 16,\n                min_holes=5,\n                min_height=config.image_size // 32,\n                min_width=config.image_size // 32,\n                fill_value=0,\n                p=0.5\n            ),  # Randomly drop large regions of the image\n            ToTensorV2(p=1.0),  # Convert image to PyTorch tensor\n        ], p=1.)\n    \n    @staticmethod\n    def get_valid_transforms(config: Config) -> Compose:\n        \"\"\"Return validation data preprocessing transforms.\"\"\"\n        return Compose([\n            Resize(config.image_size, config.image_size),  # Resize image to desired size\n            Normalize(\n                mean=config.mean,\n                std=config.std,\n                max_pixel_value=255.0,\n                p=1.0\n            ),  # Normalize image with mean and std\n            ToTensorV2(p=1.0),  # Convert image to PyTorch tensor\n        ], p=1.)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.161056Z","iopub.execute_input":"2024-11-28T19:47:44.161299Z","iopub.status.idle":"2024-11-28T19:47:44.174652Z","shell.execute_reply.started":"2024-11-28T19:47:44.161275Z","shell.execute_reply":"2024-11-28T19:47:44.173978Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ViTTrainer:\n    \"\"\"Trainer class optimized for Vision Transformer.\"\"\"\n    \n    def __init__(self, model: nn.Module, config: Config):\n        self.model = model                             # Assign the model to an instance variable\n        self.config = config                           # Store the configuration parameters\n        self.criterion = nn.CrossEntropyLoss()         # Define the loss function (Cross-Entropy Loss)\n        \n        # Optimizer with weight decay\n        self.optimizer = optim.AdamW(\n            model.parameters(),\n            lr=config.learning_rate,                   # Set learning rate from config\n            weight_decay=config.weight_decay           # Apply weight decay regularization\n        )\n        \n        # Cosine annealing scheduler\n        self.scheduler = optim.lr_scheduler.CosineAnnealingLR(\n            self.optimizer,\n            T_max=config.num_epochs,                   # Maximum number of iterations (epochs)\n            eta_min=1e-6                               # Minimum learning rate after decay\n        )\n        \n        # Initialize GradScaler for mixed precision training if CUDA is available\n        self.scaler = torch.cuda.amp.GradScaler() if config.fp16 else None\n                                                     # Use GradScaler for mixed precision if fp16 is enabled\n    \n    def train_epoch(self, train_loader: DataLoader, device: torch.device) -> float:\n        \"\"\"Train the model for one epoch.\"\"\"\n        self.model.train()                                     # Set the model to training mode\n        total_loss = 0                                         # Initialize total loss for the epoch\n        \n        with tqdm(train_loader, desc='Training') as pbar:      # Create a progress bar for the training loop\n            for batch_idx, (images, targets) in enumerate(pbar):\n                images = images.to(device)                     # Move images to the specified device (GPU or CPU)\n                targets = targets.to(device)                   # Move targets to the device\n                \n                # Mixed precision training if enabled\n                if self.config.fp16:\n                    with torch.cuda.amp.autocast():                 # Enable autocasting for mixed precision\n                        outputs = self.model(images)                # Forward pass through the model\n                        loss = self.criterion(outputs, targets)     # Compute loss between outputs and targets\n                        loss = loss / self.config.grad_accum_steps  # Normalize loss for gradient accumulation\n                    \n                    self.scaler.scale(loss).backward()         # Backward pass with scaled loss for mixed precision\n                    \n                    if (batch_idx + 1) % self.config.grad_accum_steps == 0:\n                        self.scaler.step(self.optimizer)       # Update model parameters\n                        self.scaler.update()                   # Update the scaler for next iteration\n                        self.optimizer.zero_grad()             # Reset gradients\n                else:\n                    outputs = self.model(images)               # Forward pass through the model\n                    loss = self.criterion(outputs, targets)    # Compute loss\n                    loss = loss / self.config.grad_accum_steps # Normalize loss for gradient accumulation\n                    \n                    loss.backward()                            # Backward pass\n                    \n                    if (batch_idx + 1) % self.config.grad_accum_steps == 0:\n                        self.optimizer.step()                  # Update model parameters\n                        self.optimizer.zero_grad()             # Reset gradients\n                \n                total_loss += loss.item() * self.config.grad_accum_steps  # Accumulate total loss\n                pbar.set_postfix({'loss': loss.item() * self.config.grad_accum_steps})  # Update progress bar with current loss\n        \n        self.scheduler.step()                                  # Update learning rate scheduler\n        return total_loss / len(train_loader)                  # Return average loss for the epoch\n    \n    @torch.no_grad()\n    def validate(self, valid_loader: DataLoader, device: torch.device) -> Tuple[float, float]:\n        \"\"\"Evaluate the model on the validation set.\"\"\"\n        self.model.eval()                                      # Set the model to evaluation mode\n        total_loss = 0                                         # Initialize total loss for the validation\n        predictions = []                                       # List to store predicted labels\n        targets = []                                           # List to store true labels\n        \n        with tqdm(valid_loader, desc='Validating') as pbar:    # Create a progress bar for the validation loop\n            for images, batch_targets in pbar:\n                images = images.to(device)                     # Move images to the specified device\n                batch_targets = batch_targets.to(device)       # Move targets to the device\n                \n                if self.config.fp16:\n                    with torch.cuda.amp.autocast():                    # Enable autocasting for mixed precision\n                        outputs = self.model(images)                   # Forward pass through the model\n                        loss = self.criterion(outputs, batch_targets)  # Compute loss\n                else:\n                    outputs = self.model(images)                   # Forward pass through the model\n                    loss = self.criterion(outputs, batch_targets)  # Compute loss\n                \n                total_loss += loss.item()                          # Accumulate total loss\n                predictions.extend(torch.argmax(outputs, dim=1).cpu().numpy())  # Store predicted labels\n                targets.extend(batch_targets.cpu().numpy())        # Store true labels\n                \n                pbar.set_postfix({'loss': loss.item()})            # Update progress bar with current loss\n        \n        accuracy = np.mean(np.array(predictions) == np.array(targets))  # Calculate accuracy\n        return total_loss / len(valid_loader), accuracy            # Return average loss and accuracy\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.176741Z","iopub.execute_input":"2024-11-28T19:47:44.176990Z","iopub.status.idle":"2024-11-28T19:47:44.192974Z","shell.execute_reply.started":"2024-11-28T19:47:44.176966Z","shell.execute_reply":"2024-11-28T19:47:44.192224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(config: Config, train_df: pd.DataFrame):\n    \"\"\"Train the ViT model using cross-validation.\"\"\"\n    print(f\"Using device: {config.device_type}\")  # Inform about the device being used\n    print(f\"Mixed precision training: {'enabled' if config.fp16 else 'disabled'}\")  # Inform about mixed precision\n    \n    # Create Stratified K-Fold splits to maintain class distribution\n    skf = StratifiedKFold(\n        n_splits=config.num_folds,\n        shuffle=True,\n        random_state=config.seed\n    )\n    \n    for fold, (train_idx, valid_idx) in enumerate(skf.split(train_df, train_df.label)):\n        if fold not in config.used_folds:\n            continue  # Skip folds that are not used for inference\n        \n        print(f'Training fold {fold}')  # Inform about the current fold being trained\n        \n        # Split data into training and validation sets based on indices\n        train_data = train_df.iloc[train_idx].reset_index(drop=True)\n        valid_data = train_df.iloc[valid_idx].reset_index(drop=True)\n        \n        # Create training dataset with data augmentation\n        train_dataset = CassavaDataset(\n            train_data,\n            PATHS['TRAIN_IMAGES'],\n            transforms=ViTDataTransforms.get_train_transforms(config)\n        )\n        \n        # Create validation dataset without data augmentation\n        valid_dataset = CassavaDataset(\n            valid_data,\n            PATHS['TRAIN_IMAGES'],\n            transforms=ViTDataTransforms.get_valid_transforms(config)\n        )\n        \n        # Create DataLoader for training data\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=config.train_batch_size,\n            shuffle=True,  # Shuffle training data\n            num_workers=config.num_workers,\n            pin_memory=True if config.device_type == 'cuda' else False  # Enable pin memory for CUDA\n        )\n        \n        # Create DataLoader for validation data\n        valid_loader = DataLoader(\n            valid_dataset,\n            batch_size=config.valid_batch_size,\n            shuffle=False,  # Do not shuffle validation data\n            num_workers=config.num_workers,\n            pin_memory=True if config.device_type == 'cuda' else False  # Enable pin memory for CUDA\n        )\n        \n        # Initialize the Vision Transformer model and move it to the device\n        model = CassavaViT(\n            num_classes=train_df.label.nunique(),  # Number of unique classes\n            pretrained=config.pretrained,  # Use pretrained weights\n            model_name=config.model_name  # Specify model architecture\n        ).to(config.device)\n        \n        trainer = ViTTrainer(model, config)  # Initialize the trainer\n        best_loss = float('inf')  # Initialize best loss for checkpointing\n        \n        for epoch in range(config.num_epochs):\n            print(f'Epoch {epoch + 1}/{config.num_epochs}')  # Inform about the current epoch\n            \n            train_loss = trainer.train_epoch(train_loader, config.device)  # Train for one epoch\n            valid_loss, accuracy = trainer.validate(valid_loader, config.device)  # Validate the model\n            \n            print(f'Train Loss: {train_loss:.4f}')  # Print training loss\n            print(f'Valid Loss: {valid_loss:.4f}, Accuracy: {accuracy:.4f}')  # Print validation loss and accuracy\n            \n            if valid_loss < best_loss:\n                best_loss = valid_loss  # Update best loss if current validation loss is lower\n                checkpoint_path = os.path.join(\n                    PATHS['WEIGHTS'],\n                    f'{config.model_name}_fold_{fold}_{epoch}'\n                )  # Define checkpoint path\n                torch.save(model.state_dict(), checkpoint_path)  # Save model weights\n                print(f'Saved checkpoint: {checkpoint_path}')  # Inform about saved checkpoint\n        \n        del model, trainer  # Delete model and trainer to free memory\n        if config.device_type == 'cuda':\n            torch.cuda.empty_cache()  # Clear CUDA cache to free up GPU memory\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.193921Z","iopub.execute_input":"2024-11-28T19:47:44.194230Z","iopub.status.idle":"2024-11-28T19:47:44.208473Z","shell.execute_reply.started":"2024-11-28T19:47:44.194192Z","shell.execute_reply":"2024-11-28T19:47:44.207615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"Main training function.\"\"\"\n    # Create weights directory if it doesn't exist\n    os.makedirs(PATHS['WEIGHTS'], exist_ok=True)\n    \n    config = Config()  # Initialize configuration\n    seed_everything(config.seed)  # Set random seeds for reproducibility\n    \n    # Load training data from CSV\n    train_df = pd.read_csv(PATHS['TRAIN_CSV'])\n    print(f\"Training data shape: {train_df.shape}\")  # Print shape of training data\n    \n    # Print dataset information\n    print_dataset_info(train_df)  # Display dataset statistics\n    \n    # Train the Vision Transformer model using cross-validation\n    train_model(config, train_df)\n    \n    print(\"Training completed successfully!\")  # Inform that training is complete\n\nif __name__ == '__main__':\n    main()  # Execute the main function","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T19:47:44.209429Z","iopub.execute_input":"2024-11-28T19:47:44.209668Z","iopub.status.idle":"2024-11-29T00:11:27.208143Z","shell.execute_reply.started":"2024-11-28T19:47:44.209645Z","shell.execute_reply":"2024-11-29T00:11:27.206786Z"}},"outputs":[],"execution_count":null}]}