{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":94689,"databundleVersionId":11605086,"sourceType":"competition"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Forams Classification 2025 - Final Solution\n# This notebook implements a multi-view 2D CNN approach for the semi-supervised classification of foraminifera\n\nimport os\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tifffile import imread\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score\n\n# Set random seeds for reproducibility\nnp.random.seed(42)\ntorch.manual_seed(42)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(42)\n\n# Check GPU availability\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\nif torch.cuda.is_available():\n    for i in range(torch.cuda.device_count()):\n        print(f\"GPU {i}: {torch.cuda.get_device_name(i)}\")\n\n# Define file paths\nBASE_PATH = '/kaggle/input/forams-classification-2025'\nlabeled_vols_path = f'{BASE_PATH}/volumes/volumes/labelled'\nunlabeled_vols_path = f'{BASE_PATH}/volumes/volumes/unlabelled'\nlabeled_viz_path = f'{BASE_PATH}/visualizations/visualizations/labelled'\nunlabeled_viz_path = f'{BASE_PATH}/visualizations/visualizations/unlabelled'\n\n# Read CSV files\nlabeled_df = pd.read_csv(f'{BASE_PATH}/labelled.csv')\nunlabeled_df = pd.read_csv(f'{BASE_PATH}/unlabelled.csv')\nsample_submission = pd.read_csv(f'{BASE_PATH}/sample_submission.csv')\n\nprint(\"Labeled data shape:\", labeled_df.shape)\nprint(\"Unlabeled data shape:\", unlabeled_df.shape)\nprint(\"Sample submission shape:\", sample_submission.shape)\n\n# Check label distribution\nprint(\"\\nLabel distribution:\")\nprint(labeled_df['label'].value_counts().sort_index())\n\n# Create output directory\nos.makedirs('/kaggle/working/models', exist_ok=True)\n\n# Function to visualize a volume and its slices\ndef visualize_volume(volume_path, viz_path=None, title=None):\n    # Load the volume\n    volume = imread(volume_path)\n    \n    # Get middle slices in each dimension\n    slice_x = volume[64, :, :]\n    slice_y = volume[:, 64, :]\n    slice_z = volume[:, :, 64]\n    \n    # Create a figure with subplots\n    fig, axes = plt.subplots(1, 4, figsize=(20, 5))\n    \n    # Plot slices\n    axes[0].imshow(slice_x, cmap='gray')\n    axes[0].set_title('X-Slice (64)')\n    \n    axes[1].imshow(slice_y, cmap='gray')\n    axes[1].set_title('Y-Slice (64)')\n    \n    axes[2].imshow(slice_z, cmap='gray')\n    axes[2].set_title('Z-Slice (64)')\n    \n    # Plot the visualization if provided\n    if viz_path and os.path.exists(viz_path):\n        viz_img = plt.imread(viz_path)\n        axes[3].imshow(viz_img)\n        axes[3].set_title('Visualization')\n    else:\n        axes[3].axis('off')\n    \n    if title:\n        plt.suptitle(title)\n    plt.tight_layout()\n    plt.show()\n\n# Multi-view Dataset class that extracts key 2D slices from the 3D volumes\nclass ForamMultiViewDataset(Dataset):\n    def __init__(self, file_paths, labels=None, transform=None, n_slices=3):\n        self.file_paths = file_paths\n        self.labels = labels  # None for unlabeled data\n        self.transform = transform\n        self.n_slices = n_slices  # Number of slices per dimension\n        \n    def __len__(self):\n        return len(self.file_paths)\n    \n    def __getitem__(self, idx):\n        # Load the 3D volume\n        volume_path = self.file_paths[idx]\n        volume = imread(volume_path).astype(np.float32)\n        \n        # Select key slices from each dimension\n        slices = []\n        \n        # Get center slices and slices around 1/4 and 3/4 of each dimension\n        dim_size = volume.shape[0]  # Assuming cubic volume\n        indices = [dim_size // 4, dim_size // 2, 3 * dim_size // 4]\n        \n        # Extract slices from different planes (axial, coronal, sagittal)\n        for i in indices:\n            slices.append(volume[i, :, :])  # xy plane (axial)\n            slices.append(volume[:, i, :])  # xz plane (coronal)\n            slices.append(volume[:, :, i])  # yz plane (sagittal)\n        \n        # Convert to tensor and normalize\n        slices = [torch.from_numpy(slice).float() / 255.0 for slice in slices]\n        \n        # Apply transforms if any\n        if self.transform:\n            slices = [self.transform(slice.unsqueeze(0)).squeeze(0) for slice in slices]\n        \n        # Stack slices as channels\n        multi_view = torch.stack(slices)\n        \n        # Return volume and label (if available)\n        if self.labels is not None:\n            return multi_view, self.labels[idx]\n        else:\n            return multi_view\n\n# Data augmentation classes\nclass RandomBrightness:\n    def __init__(self, factor=0.2):\n        self.factor = factor\n        \n    def __call__(self, x):\n        factor = np.random.uniform(1.0 - self.factor, 1.0 + self.factor)\n        x = x * factor\n        return torch.clamp(x, 0, 1)\n\nclass RandomContrast:\n    def __init__(self, factor=0.2):\n        self.factor = factor\n        \n    def __call__(self, x):\n        factor = np.random.uniform(1.0 - self.factor, 1.0 + self.factor)\n        mean = x.mean()\n        x = (x - mean) * factor + mean\n        return torch.clamp(x, 0, 1)\n\nclass RandomGamma:\n    def __init__(self, range=(0.7, 1.5)):\n        self.range = range\n        \n    def __call__(self, x):\n        gamma = np.random.uniform(self.range[0], self.range[1])\n        return torch.pow(x, gamma)\n\nclass Compose:\n    def __init__(self, transforms):\n        self.transforms = transforms\n        \n    def __call__(self, x):\n        for t in self.transforms:\n            x = t(x)\n        return x\n\n# Create transform compositions\ntrain_transform = Compose([\n    RandomBrightness(factor=0.3),\n    RandomContrast(factor=0.3),\n    RandomGamma(range=(0.7, 1.3))\n])\n\n# No transforms for validation and test\nval_transform = None\n\n# 2D CNN model for multi-view processing\nclass MultiViewForamCNN(nn.Module):\n    def __init__(self, num_classes=15, num_views=9):\n        super(MultiViewForamCNN, self).__init__()\n        \n        # Input: num_views x 128 x 128\n        self.conv1 = nn.Conv2d(num_views, 32, kernel_size=3, padding=1)\n        self.bn1 = nn.BatchNorm2d(32)\n        self.pool1 = nn.MaxPool2d(kernel_size=2)\n        \n        # After pool1: 32 x 64 x 64\n        self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)\n        self.bn2 = nn.BatchNorm2d(64)\n        self.pool2 = nn.MaxPool2d(kernel_size=2)\n        \n        # After pool2: 64 x 32 x 32\n        self.conv3 = nn.Conv2d(64, 128, kernel_size=3, padding=1)\n        self.bn3 = nn.BatchNorm2d(128)\n        self.pool3 = nn.MaxPool2d(kernel_size=2)\n        \n        # After pool3: 128 x 16 x 16\n        self.conv4 = nn.Conv2d(128, 256, kernel_size=3, padding=1)\n        self.bn4 = nn.BatchNorm2d(256)\n        self.pool4 = nn.MaxPool2d(kernel_size=2)\n        \n        # After pool4: 256 x 8 x 8\n        self.global_avg_pool = nn.AdaptiveAvgPool2d(1)\n        \n        # After global_avg_pool: 256 x 1 x 1\n        self.fc1 = nn.Linear(256, 128)\n        self.dropout = nn.Dropout(0.5)\n        self.fc2 = nn.Linear(128, num_classes)\n        \n    def forward(self, x):\n        # Feature extraction\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.pool1(x)\n        \n        x = F.relu(self.bn2(self.conv2(x)))\n        x = self.pool2(x)\n        \n        x = F.relu(self.bn3(self.conv3(x)))\n        x = self.pool3(x)\n        \n        x = F.relu(self.bn4(self.conv4(x)))\n        x = self.pool4(x)\n        \n        # Global average pooling\n        x = self.global_avg_pool(x)\n        x = x.view(x.size(0), -1)\n        \n        # Classification\n        features = F.relu(self.fc1(x))\n        x = self.dropout(features)\n        logits = self.fc2(x)\n        \n        return logits, features\n\n# Prepare file paths and labels for datasets\ndef prepare_datasets(labeled_vols_path, unlabeled_vols_path, labeled_df, val_size=0.2):\n    # Get file paths for labeled data\n    labeled_files = []\n    labeled_labels = []\n    \n    for _, row in labeled_df.iterrows():\n        id_num = row['id'].split('_')[-1]\n        matching_files = [f for f in os.listdir(labeled_vols_path) if f.startswith(f\"labelled_foram_{id_num}_\")]\n        if matching_files:\n            labeled_files.append(os.path.join(labeled_vols_path, matching_files[0]))\n            labeled_labels.append(row['label'])\n    \n    # Split labeled data into train and validation\n    train_files, val_files, train_labels, val_labels = train_test_split(\n        labeled_files, labeled_labels, test_size=val_size, stratify=labeled_labels, random_state=42\n    )\n    \n    # Get file paths for unlabeled data\n    unlabeled_files = [os.path.join(unlabeled_vols_path, f) for f in os.listdir(unlabeled_vols_path) \n                       if f.endswith('.tif')]\n    \n    return train_files, train_labels, val_files, val_labels, unlabeled_files\n\n# Prepare datasets\ntrain_files, train_labels, val_files, val_labels, unlabeled_files = prepare_datasets(\n    labeled_vols_path, unlabeled_vols_path, labeled_df\n)\n\nprint(f\"Training samples: {len(train_files)}\")\nprint(f\"Validation samples: {len(val_files)}\")\nprint(f\"Unlabeled samples: {len(unlabeled_files)}\")\n\n# Create datasets\ntrain_dataset = ForamMultiViewDataset(train_files, train_labels, transform=train_transform)\nval_dataset = ForamMultiViewDataset(val_files, val_labels, transform=val_transform)\n\n# Create dataloaders\nbatch_size = 8\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=2)\n\n# Initialize model\nmodel = MultiViewForamCNN(num_classes=15, num_views=9)\nmodel = model.to(device)\n\n# Define loss function and optimizer\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-5)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=5)\n\n# Training function\ndef train_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    for inputs, targets in loader:\n        inputs, targets = inputs.to(device), targets.to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs, _ = model(inputs)\n        loss = criterion(outputs, targets)\n        \n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * inputs.size(0)\n        _, predicted = outputs.max(1)\n        total += targets.size(0)\n        correct += predicted.eq(targets).sum().item()\n    \n    epoch_loss = running_loss / total\n    epoch_acc = correct / total\n    \n    return epoch_loss, epoch_acc\n\n# Evaluation function\ndef evaluate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    all_targets = []\n    all_preds = []\n    \n    with torch.no_grad():\n        for inputs, targets in loader:\n            inputs, targets = inputs.to(device), targets.to(device)\n            \n            outputs, _ = model(inputs)\n            loss = criterion(outputs, targets)\n            \n            running_loss += loss.item() * inputs.size(0)\n            _, predicted = outputs.max(1)\n            total += targets.size(0)\n            correct += predicted.eq(targets).sum().item()\n            \n            all_targets.extend(targets.cpu().numpy())\n            all_preds.extend(predicted.cpu().numpy())\n    \n    epoch_loss = running_loss / total\n    epoch_acc = correct / total\n    \n    f1 = f1_score(all_targets, all_preds, average='macro')\n    \n    return epoch_loss, epoch_acc, f1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T09:15:53.614330Z","iopub.execute_input":"2025-04-04T09:15:53.614640Z","iopub.status.idle":"2025-04-04T09:15:53.822628Z","shell.execute_reply.started":"2025-04-04T09:15:53.614614Z","shell.execute_reply":"2025-04-04T09:15:53.821819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Train the model\nnum_epochs = 150  # Train for more epochs for final submission\nbest_f1 = 0.0\nbest_model_path = '/kaggle/working/models/best_model.pt'\n\nprint(\"Starting training with multi-view 2D CNN...\")\nfor epoch in range(num_epochs):\n    train_loss, train_acc = train_epoch(model, train_loader, optimizer, criterion, device)\n    val_loss, val_acc, val_f1 = evaluate(model, val_loader, criterion, device)\n    \n    # Print progress\n    print(f\"Epoch {epoch+1}/{num_epochs}: \"\n          f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}, \"\n          f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}, Val F1: {val_f1:.4f}\")\n    \n    # Update learning rate\n    scheduler.step(val_f1)\n    \n    # Save best model\n    if val_f1 > best_f1:\n        best_f1 = val_f1\n        torch.save(model.state_dict(), best_model_path)\n        print(f\"Saved new best model with F1: {best_f1:.4f}\")\n\n# Create dataset for unlabeled data\nprint(\"Creating dataset for unlabeled data...\")\nunlabeled_dataset = ForamMultiViewDataset(unlabeled_files, transform=val_transform)\nunlabeled_loader = DataLoader(unlabeled_dataset, batch_size=16, shuffle=False, num_workers=2)\n\n# Load the best model\nprint(\"Loading best model for prediction...\")\nmodel.load_state_dict(torch.load(best_model_path))\nmodel.eval()\n\n# Generate predictions for unlabeled data\nprint(\"Generating predictions for unlabeled data...\")\nunlabeled_preds = []\nunlabeled_probs = []\nunlabeled_ids = []\n\nwith torch.no_grad():\n    for i, batch in enumerate(unlabeled_loader):\n        if i % 100 == 0:\n            print(f\"Processing batch {i}/{len(unlabeled_loader)}\")\n        \n        # Get file paths for the current batch\n        batch_files = unlabeled_files[i*unlabeled_loader.batch_size:min((i+1)*unlabeled_loader.batch_size, len(unlabeled_files))]\n        batch_ids = [int(os.path.basename(f).split('_')[1]) for f in batch_files]\n        \n        # Forward pass\n        inputs = batch.to(device)\n        outputs, _ = model(inputs)\n        \n        # Get predictions and probabilities\n        probs = F.softmax(outputs, dim=1)\n        max_probs, preds = probs.max(1)\n        \n        # Store predictions, probabilities, and IDs\n        unlabeled_preds.extend(preds.cpu().numpy())\n        unlabeled_probs.extend(max_probs.cpu().numpy())\n        unlabeled_ids.extend(batch_ids)\n\n# Label any low-confidence predictions as 'unknown' (class 14)\nconfidence_threshold = 0.8\nfor i in range(len(unlabeled_probs)):\n    if unlabeled_probs[i] < confidence_threshold:\n        unlabeled_preds[i] = 14  # Assign to unknown class\n\n# Create submission dataframe\nsubmission_df = pd.DataFrame({'id': unlabeled_ids, 'label': unlabeled_preds})\n\n# Sort by ID to match the order in sample_submission\nsubmission_df = submission_df.sort_values('id').reset_index(drop=True)\n\n# Save submission file\nsubmission_path = '/kaggle/working/submission.csv'\nsubmission_df.to_csv(submission_path, index=False)\nprint(f\"Submission saved to {submission_path}\")\n\n# Display submission statistics\nprint(\"Prediction distribution:\")\nprint(submission_df['label'].value_counts().sort_index())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T09:16:27.878066Z","iopub.execute_input":"2025-04-04T09:16:27.878361Z","iopub.status.idle":"2025-04-04T09:20:57.335736Z","shell.execute_reply.started":"2025-04-04T09:16:27.878338Z","shell.execute_reply":"2025-04-04T09:20:57.334845Z"}},"outputs":[],"execution_count":null}]}