{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"********Configuration de l'environnement et Accès aux données********","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\nfrom sklearn.model_selection import train_test_split\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n\n!pip install timm -q\nimport timm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Utilisation de l'appareil : {device}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2026-04-18T17:38:32.117Z","iopub.execute_input":"2026-04-18T17:38:32.117691Z","iopub.status.idle":"2026-04-18T17:38:35.459098Z","shell.execute_reply.started":"2026-04-18T17:38:32.117657Z","shell.execute_reply":"2026-04-18T17:38:35.4581Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/input/competitions/aptos2019-blindness-detection/train.csv')\nprint(f\"Nombre total d'images : {len(train_df)}\")\nprint(train_df['diagnosis'].value_counts())\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T17:40:10.953415Z","iopub.execute_input":"2026-04-18T17:40:10.954143Z","iopub.status.idle":"2026-04-18T17:40:10.999825Z","shell.execute_reply.started":"2026-04-18T17:40:10.954113Z","shell.execute_reply":"2026-04-18T17:40:10.999134Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"****Prétraitement et Création du Dataset PyTorch****","metadata":{}},{"cell_type":"code","source":"# 1. Définition des paramètres\nIMG_SIZE = 224\nBATCH_SIZE = 32\n\n# 2. Split des données (80% train, 20% validation)\ntrain_data, val_data = train_test_split(\n    train_df, \n    test_size=0.2, \n    random_state=42, \n    stratify=train_df['diagnosis']\n)\n\n# 3. Transformations (Augmentation pour le train, simple redimensionnement pour le val)\ntrain_transforms = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(20),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) \n])\n\nval_transforms = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\n# 4. Classe Dataset personnalisée\nclass APTOSDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_name = os.path.join(self.img_dir, self.df.iloc[idx, 0] + '.png')\n        image = Image.open(img_name).convert('RGB')\n        label = self.df.iloc[idx, 1]\n        \n        if self.transform:\n            image = self.transform(image)\n            \n        return image, torch.tensor(label, dtype=torch.long)\n\n# 5. Création des DataLoaders\ntrain_dataset = APTOSDataset(train_data, '/kaggle/input/competitions/aptos2019-blindness-detection/train_images', train_transforms)\nval_dataset = APTOSDataset(val_data, '/kaggle/input/competitions/aptos2019-blindness-detection/train_images', val_transforms)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False)\n\nprint(f\"Taille du set d'entraînement : {len(train_dataset)}\")\nprint(f\"Taille du set de validation : {len(val_dataset)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T17:42:10.814176Z","iopub.execute_input":"2026-04-18T17:42:10.814987Z","iopub.status.idle":"2026-04-18T17:42:10.835389Z","shell.execute_reply.started":"2026-04-18T17:42:10.814954Z","shell.execute_reply":"2026-04-18T17:42:10.83435Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"****Définition du modèle Vision Transformer (ViT)****","metadata":{}},{"cell_type":"code","source":"# 1. Création du modèle ViT\nmodel_name = 'vit_base_patch16_224'\nmodel = timm.create_model(model_name, pretrained=True)\n\n# 2. Adaptation de la \"tête\" de classification\nn_features = model.head.in_features\nmodel.head = nn.Linear(n_features, 5)\n\nmodel = model.to(device)\n\n# 3. Définition de la Loss et de l'Optimiseur\ncriterion = nn.CrossEntropyLoss()\n\n# AdamW est généralement plus performant que Adam pour les Transformers\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5)\n\n# Planificateur de taux d'apprentissage (pour réduire le LR si la perte stagne)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=2)\n\nprint(f\"Modèle {model_name} chargé et adapté avec succès.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T17:43:22.49524Z","iopub.execute_input":"2026-04-18T17:43:22.495714Z","iopub.status.idle":"2026-04-18T17:43:28.07343Z","shell.execute_reply.started":"2026-04-18T17:43:22.495687Z","shell.execute_reply":"2026-04-18T17:43:28.072533Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"****La Boucle d'Entraînement et de Validation****","metadata":{}},{"cell_type":"code","source":"import time\nfrom tqdm import tqdm\n\ndef train_one_epoch(model, dataloader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    for images, labels in tqdm(dataloader, desc=\"Training\"):\n        images, labels = images.to(device), labels.to(device)\n        \n        # Reset des gradients\n        optimizer.zero_grad()\n        \n        # Forward pass\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        \n        # Backward pass et optimisation\n        loss.backward()\n        optimizer.step()\n        \n        # Statistiques\n        running_loss += loss.item() * images.size(0)\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n    return running_loss / total, 100. * correct / total\n\ndef validate(model, dataloader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad(): # Pas de calcul de gradient en validation\n        for images, labels in tqdm(dataloader, desc=\"Validating\"):\n            images, labels = images.to(device), labels.to(device)\n            \n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item() * images.size(0)\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n            \n    return running_loss / total, 100. * correct / total\n\n# --- Lancement de l'entraînement ---\nnum_epochs = 10\nbest_val_acc = 0.0\n\nfor epoch in range(num_epochs):\n    print(f\"\\nÉpoque {epoch+1}/{num_epochs}\")\n    \n    train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)\n    val_loss, val_acc = validate(model, val_loader, criterion, device)\n    \n    scheduler.step(val_loss)\n    \n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}%\")\n    print(f\"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.2f}%\")\n    \n    # Sauvegarder le meilleur modèle pour la phase d'interprétabilité\n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        torch.save(model.state_dict(), 'best_vit_aptos.pth')\n        print(\"Modèle sauvegardé !\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T17:44:52.400184Z","iopub.execute_input":"2026-04-18T17:44:52.400886Z","iopub.status.idle":"2026-04-18T19:11:34.758762Z","shell.execute_reply.started":"2026-04-18T17:44:52.400856Z","shell.execute_reply":"2026-04-18T19:11:34.758082Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"****Évaluation Approfondie et Matrice de Confusion****","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, classification_report, cohen_kappa_score\nimport seaborn as sns\n\n# 1. Charger les meilleurs poids sauvegardés\nmodel.load_state_dict(torch.load('best_vit_aptos.pth'))\nmodel.eval()\n\nall_preds = []\nall_labels = []\n\n# 2. Prédiction sur le set de validation\nwith torch.no_grad():\n    for images, labels in tqdm(val_loader, desc=\"Evaluation\"):\n        images = images.to(device)\n        outputs = model(images)\n        _, preds = torch.max(outputs, 1)\n        \n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.numpy())\n\n\nkappa = cohen_kappa_score(all_labels, all_preds, weights='quadratic')\nprint(f\"\\nScore Quadratic Weighted Kappa : {kappa:.4f}\")\n\n\nprint(\"\\nClassification Report :\")\nprint(classification_report(all_labels, all_preds, target_names=['0', '1', '2', '3', '4']))\n\ncm = confusion_matrix(all_labels, all_preds)\nplt.figure(figsize=(10, 8))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=['0', '1', '2', '3', '4'], yticklabels=['0', '1', '2', '3', '4'])\nplt.xlabel('Prédictions')\nplt.ylabel('Vrais Labels')\nplt.title('Matrice de Confusion - ViT sur APTOS')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T19:13:57.865Z","iopub.execute_input":"2026-04-18T19:13:57.865478Z","iopub.status.idle":"2026-04-18T19:15:26.372351Z","shell.execute_reply.started":"2026-04-18T19:13:57.865448Z","shell.execute_reply":"2026-04-18T19:15:26.371714Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"****Interprétabilité - Visualisation des \"Attention Maps\"****","metadata":{}},{"cell_type":"code","source":"import cv2\n\ndef visualize_attention(model, dataset, idx, device):\n    model.eval()\n    \n    image_tensor, label = dataset[idx]\n    image_input = image_tensor.unsqueeze(0).to(device)\n    \n\n    attentions = []\n    def hook_fn(module, input, output):\n        attentions.append(output)\n\n    handle = model.blocks[-1].attn.attn_drop.register_forward_hook(hook_fn)\n    \n    # 3. Forward pass\n    with torch.no_grad():\n        output = model(image_input)\n        _, pred = torch.max(output, 1)\n        \n        # Extraction des features pour la heatmap\n        feature_map = model.forward_features(image_input) \n        feature_map = feature_map[:, 1:, :] \n\n    handle.remove()\n\n    heatmap = feature_map.abs().mean(-1).reshape(14, 14).detach().cpu().numpy()\n    \n    # 5. Affichage et Dé-normalisation\n    img_display = image_tensor.permute(1, 2, 0).cpu().numpy()\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    img_display = std * img_display + mean\n    img_display = np.clip(img_display, 0, 1)\n\n    plt.figure(figsize=(12, 5))\n    \n    plt.subplot(1, 2, 1)\n    plt.imshow(img_display)\n    plt.title(f\"Originale (Vrai: {label}, Prédit: {pred.item()})\")\n    plt.axis('off')\n\n    plt.subplot(1, 2, 2)\n    plt.imshow(img_display)\n    heatmap_resized = cv2.resize(heatmap, (224, 224))\n    plt.imshow(heatmap_resized, alpha=0.5, cmap='jet')\n    plt.title(\"Attention du ViT (Dernière couche)\")\n    plt.axis('off')\n    \n    plt.show()\n\n\nvisualize_attention(model, val_dataset, idx=40, device=device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T19:19:42.281841Z","iopub.execute_input":"2026-04-18T19:19:42.2826Z","iopub.status.idle":"2026-04-18T19:19:42.856776Z","shell.execute_reply.started":"2026-04-18T19:19:42.282571Z","shell.execute_reply":"2026-04-18T19:19:42.855985Z"}},"outputs":[],"execution_count":null}]}