{"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":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":6917177,"sourceType":"datasetVersion","datasetId":3889865},{"sourceId":7001173,"sourceType":"datasetVersion","datasetId":3881967},{"sourceId":7029230,"sourceType":"datasetVersion","datasetId":4027203}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"**Chargement dataset**","metadata":{}},{"cell_type":"code","source":"import os, glob\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nDATASET_FOLDER = \"/kaggle/input/UBC-OCEAN/\"\nDATASET_TILES = \"/kaggle/input/ubc-ocean-tiles-w-masks-2048px-scale-0-25\"\nDATASET_IMAGES = os.path.join(DATASET_TILES, \"train_images\")\nDATASET_MASKS = os.path.join(DATASET_TILES, \"train_masks\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:03:13.040049Z","iopub.execute_input":"2025-05-29T18:03:13.040382Z","iopub.status.idle":"2025-05-29T18:03:13.33901Z","shell.execute_reply.started":"2025-05-29T18:03:13.040354Z","shell.execute_reply":"2025-05-29T18:03:13.33841Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Nettoyage initial du DataFrame**","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train.csv\"))\nlabels = list(df_train[\"label\"].unique())\nprint(f\"Dataset/train size: {len(df_train)}\")\ndisplay(df_train.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:03:13.340111Z","iopub.execute_input":"2025-05-29T18:03:13.340575Z","iopub.status.idle":"2025-05-29T18:03:13.375612Z","shell.execute_reply.started":"2025-05-29T18:03:13.340542Z","shell.execute_reply":"2025-05-29T18:03:13.374811Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Analyse exploratoire des classes**","metadata":{}},{"cell_type":"code","source":"inds = list(map(os.path.basename, glob.glob(os.path.join(DATASET_IMAGES, \"*\"))))\nprint(f\"{inds=}\")\ndf_train = df_train[df_train[\"image_id\"].astype(str).isin(inds)]\nprint(f\"Dataset/train size: {len(df_train)}\")\ndisplay(df_train.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:03:13.377171Z","iopub.execute_input":"2025-05-29T18:03:13.377451Z","iopub.status.idle":"2025-05-29T18:03:13.412746Z","shell.execute_reply.started":"2025-05-29T18:03:13.377431Z","shell.execute_reply":"2025-05-29T18:03:13.412119Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\n# Exemple : df_data contient une colonne 'label' avec les classes\nnum_classes = df_train['label'].nunique()\nprint(f\"Nombre de classes : {num_classes}\")\n\n# Voir aussi la liste des classes uniques\nclasses = df_train['label'].unique()\nprint(f\"Classes : {classes}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:03:13.413898Z","iopub.execute_input":"2025-05-29T18:03:13.41413Z","iopub.status.idle":"2025-05-29T18:03:13.41904Z","shell.execute_reply.started":"2025-05-29T18:03:13.414111Z","shell.execute_reply":"2025-05-29T18:03:13.418401Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"_= df_train[[\"label\"]].value_counts().plot.pie(autopct='%1.1f%%', ylabel=\"label\", figsize=(3,3))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:03:13.419848Z","iopub.execute_input":"2025-05-29T18:03:13.420122Z","iopub.status.idle":"2025-05-29T18:03:13.641282Z","shell.execute_reply.started":"2025-05-29T18:03:13.420092Z","shell.execute_reply":"2025-05-29T18:03:13.640537Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Visualisation d'exemples d'images**","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nnb_samples = 5\n\nn, m = len(np.unique(df_train['label'])), nb_samples\nfig, axarr = plt.subplots(nrows=n, ncols=m, figsize=(m * 2, n * 2))\nfor ilb, (lb, df_) in enumerate(df_train.groupby('label')):\n    img_ids = list(df_['image_id'])\n    for i in range(m):\n        if i == 0:\n            axarr[ilb, i].set_title(f\"{lb} #{len(df_)}\")\n        ls_imgs = glob.glob(os.path.join(DATASET_IMAGES, str(img_ids[i]), \"*.png\"))\n        img_path = ls_imgs[0]\n        img = plt.imread(img_path)\n        mask = np.sum(img[..., :3], axis=2) == 0\n        img[mask, :] = 255\n        axarr[ilb, i].imshow(img)\n        # axarr[ilb, i].set_xticks([])\n        # axarr[ilb, i].set_yticks([])\n_= plt.axis('off')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:03:13.642028Z","iopub.execute_input":"2025-05-29T18:03:13.642227Z","iopub.status.idle":"2025-05-29T18:03:19.1505Z","shell.execute_reply.started":"2025-05-29T18:03:13.642209Z","shell.execute_reply":"2025-05-29T18:03:19.14944Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Calcul des statistiques de couleur**","metadata":{}},{"cell_type":"code","source":"from PIL import Image\nfrom joblib import Parallel, delayed\nfrom tqdm.auto import tqdm\n\ndef _color_means(img_path):\n    img = np.array(Image.open(img_path))\n    mask = np.sum(img[..., :3], axis=2) == 0\n    img[mask, :] = 255\n    if np.max(img) > 1.5:\n        img = img / 255.0\n    clr_mean = {i: np.mean(img[..., i]) for i in range(3)}\n    clr_std = {i: np.std(img[..., i]) for i in range(3)}\n    return clr_mean, clr_std\n\n# os.path.join(DATASET_SMALL_FOLDER, \"train_images\")\nls_images = glob.glob(os.path.join(DATASET_IMAGES, \"*\", \"*.png\"))\nclr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(ls_images[:10_000]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:03:19.151214Z","iopub.execute_input":"2025-05-29T18:03:19.151569Z","iopub.status.idle":"2025-05-29T18:05:55.654723Z","shell.execute_reply.started":"2025-05-29T18:03:19.151532Z","shell.execute_reply":"2025-05-29T18:05:55.653994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_color_mean = pd.DataFrame([c[0] for c in clr_mean_std]).describe()\ndisplay(img_color_mean.T)\nimg_color_std = pd.DataFrame([c[1] for c in clr_mean_std]).describe()\ndisplay(img_color_std.T)\n\nimg_color_mean = list(img_color_mean.T[\"mean\"])\nimg_color_std = list(img_color_std.T[\"mean\"])\nprint(f\"{img_color_mean=}\\n{img_color_std=}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:05:55.656953Z","iopub.execute_input":"2025-05-29T18:05:55.65719Z","iopub.status.idle":"2025-05-29T18:05:55.714447Z","shell.execute_reply.started":"2025-05-29T18:05:55.657169Z","shell.execute_reply":"2025-05-29T18:05:55.713817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -U -q pytorch-lightning","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:05:55.715583Z","iopub.execute_input":"2025-05-29T18:05:55.715803Z","iopub.status.idle":"2025-05-29T18:06:00.934775Z","shell.execute_reply.started":"2025-05-29T18:05:55.715783Z","shell.execute_reply":"2025-05-29T18:06:00.933992Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Définition du Dataset personnalisé**","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport random\nimport numpy as np\nimport torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass CancerTilesDataset(Dataset):\n    def __init__(\n        self,\n        df_data,\n        path_img_dir: str,\n        path_mask_dir: str,\n        transforms=None,\n        split: float = 0.9,\n        mode: str = 'train',\n        labels_lut=None,\n        one_hot: bool = True,\n        tumor_thr: float = 0.3,\n    ):\n        assert os.path.isdir(path_img_dir), f\"Dossier image introuvable : {path_img_dir}\"\n        assert os.path.isdir(path_mask_dir), f\"Dossier masque introuvable : {path_mask_dir}\"\n\n        self.path_img_dir = path_img_dir\n        self.path_mask_dir = path_mask_dir\n        self.transforms = transforms\n        self.split = split\n        self.mode = mode\n        self.tumor_thr = tumor_thr\n        self.one_hot = one_hot\n        self.data = df_data\n\n        # Génération ou utilisation du LUT\n        self.labels_unique = sorted(df_data[\"label\"].unique()) + [\"Other\"]\n        self.labels_lut = labels_lut or {lb: i for i, lb in enumerate(self.labels_unique)}\n\n        # Chargement des chemins vers toutes les images\n        ls_img = sorted(glob.glob(os.path.join(self.path_img_dir, \"*\", \"*.png\")))\n        random.Random(42).shuffle(ls_img)\n        self.imgs = [(os.path.basename(os.path.dirname(p)), os.path.basename(p)) for p in ls_img]\n\n        # Split en train / val\n        assert 0.0 <= self.split <= 1.0\n        frac = int(self.split * len(self.imgs))\n        self.imgs = self.imgs[:frac] if mode == 'train' else self.imgs[frac:]\n\n    @property\n    def num_classes(self) -> int:\n        return len(self.labels_lut)\n\n    def to_one_hot(self, label: str) -> torch.Tensor:\n        one_hot = torch.zeros(self.num_classes, dtype=torch.float32)\n        index = self.labels_lut.get(label, self.labels_lut[\"Other\"])\n        one_hot[index] = 1.0\n        return one_hot\n\n    def __getitem__(self, idx: int) -> tuple:\n        image_id, tile = self.imgs[idx]\n        img_path = os.path.join(self.path_img_dir, image_id, tile)\n        mask_path = os.path.join(self.path_mask_dir, image_id, tile)\n\n        # Chargement image\n        img = np.array(Image.open(img_path))[..., :3]\n        black_bg = np.sum(img, axis=2) == 0\n        img[black_bg] = 255\n\n        # Chargement masque\n        mask = np.array(Image.open(mask_path))[..., :3]\n        tumor_mask_ratio = float(np.sum(mask == 1)) / np.prod(mask.shape)\n\n        # Récupération du label\n        row = self.data[self.data[\"image_id\"] == int(image_id)]\n        if len(row) > 0:\n            label = row.iloc[0][\"label\"]\n        else:\n            label = \"Other\"\n\n        if tumor_mask_ratio < self.tumor_thr:\n            label = \"Other\"\n\n        # Encodage\n        if self.one_hot:\n            label_tensor = self.to_one_hot(label)\n        else:\n            label_tensor = torch.tensor(self.labels_lut.get(label, self.labels_lut[\"Other\"]), dtype=torch.long)\n\n        # Transfo\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n        else:\n            img = torch.from_numpy(img).permute(2, 0, 1).float() / 255.\n\n        return img, label_tensor\n\n    def __len__(self) -> int:\n        return len(self.imgs)\ndataset = CancerTilesDataset(df_train, DATASET_IMAGES, DATASET_MASKS, one_hot=False)\n\n# Visualisation rapide\nimport matplotlib.pyplot as plt\n\nfig, axes = plt.subplots(nrows=3, ncols=3, figsize=(10, 10))\nfor i in range(9):\n    img, lb = dataset[i]\n    img_np = img.permute(1, 2, 0).numpy() if isinstance(img, torch.Tensor) else img\n    ax = axes[i // 3, i % 3]\n    ax.imshow(img_np)\n    ax.set_title(f\"{lb}\\nDims: {img.shape}\")\n    ax.axis('off')\nfig.tight_layout()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:06:00.935859Z","iopub.execute_input":"2025-05-29T18:06:00.936119Z","iopub.status.idle":"2025-05-29T18:06:06.333211Z","shell.execute_reply.started":"2025-05-29T18:06:00.936095Z","shell.execute_reply":"2025-05-29T18:06:06.332294Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Transformations**","metadata":{}},{"cell_type":"code","source":"from torchvision import transforms\n\n# Valeurs à ajuster selon ton dataset\nimg_color_mean = [0.702, 0.546, 0.696]\nimg_color_std = [0.230, 0.278, 0.195]\n\nTRAIN_TRANSFORM = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(degrees=20),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=img_color_mean, std=img_color_std)\n])\n\nVALID_TRANSFORM = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=img_color_mean, std=img_color_std)\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:06:06.334162Z","iopub.execute_input":"2025-05-29T18:06:06.334545Z","iopub.status.idle":"2025-05-29T18:06:08.669088Z","shell.execute_reply.started":"2025-05-29T18:06:06.334521Z","shell.execute_reply":"2025-05-29T18:06:08.668459Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import operator\nimport torchvision\nfrom lightning_utilities import compare_version\n\nTV_GE_0_16 = compare_version(\"torchvision\", operator.ge, \"0.16.0\")\nprint(TV_GE_0_16)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:06:08.669956Z","iopub.execute_input":"2025-05-29T18:06:08.670475Z","iopub.status.idle":"2025-05-29T18:06:08.690778Z","shell.execute_reply.started":"2025-05-29T18:06:08.670439Z","shell.execute_reply":"2025-05-29T18:06:08.690125Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**DataModule PyTorch Lightning**","metadata":{}},{"cell_type":"code","source":"import multiprocessing as mproc\nimport pytorch_lightning as pl\nfrom torch.utils.data import DataLoader\n\nclass CancerSubtypeDM(pl.LightningDataModule):\n\n    def __init__(\n        self,\n        df_data,\n        path_img_dir: str,\n        path_mask_dir: str,\n        batch_size: int = 32,\n        num_workers: int = None,\n        train_transforms = TRAIN_TRANSFORM,\n        valid_transforms = VALID_TRANSFORM,\n        one_hot: bool = True,\n    ):\n        super().__init__()\n        self.df_data = df_data\n        self.path_img_dir = path_img_dir\n        self.path_mask_dir = path_mask_dir\n        self.batch_size = batch_size\n        self.num_workers = num_workers or mproc.cpu_count()\n        self.train_dataset = None\n        self.valid_dataset = None\n        self.train_transforms = train_transforms\n        self.valid_transforms = valid_transforms\n        self.one_hot = one_hot\n\n    def prepare_data(self):\n        pass\n\n    @property\n    def num_classes(self) -> int:\n        assert self.train_dataset and self.valid_dataset\n        return len(set(self.train_dataset.labels_unique + self.valid_dataset.labels_unique))\n\n    def setup(self, stage=None):\n        self.train_dataset = CancerTilesDataset(\n            self.df_data, self.path_img_dir, self.path_mask_dir,\n            mode='train', transforms=self.train_transforms, one_hot=self.one_hot)\n        print(f\"training dataset: {len(self.train_dataset)}\")\n        self.valid_dataset = CancerTilesDataset(\n            self.df_data, self.path_img_dir, self.path_mask_dir,\n            mode='valid', transforms=self.valid_transforms, one_hot=self.one_hot,\n            # as validation is subsampled it may happen that some labels are missing\n            # and so created one-hot-encoding vector will be sorter\n            labels_lut=self.train_dataset.labels_lut)\n        print(f\"validation dataset: {len(self.valid_dataset)}\")\n\n    def train_dataloader(self):\n        return DataLoader(\n            self.train_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=True,\n        )\n\n    def val_dataloader(self):\n        return DataLoader(\n            self.valid_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=False,\n        )\n\n    def test_dataloader(self):\n        pass\n\n# ==============================\n# ==============================\n\ndm = CancerSubtypeDM(\n    df_train, DATASET_IMAGES, DATASET_MASKS,\n    batch_size=18, one_hot=not TV_GE_0_16\n)\ndm.setup()\nprint(dm.num_classes)\n\n# quick view\nfig = plt.figure(figsize=(3, 9))\nfor imgs, lbs in dm.train_dataloader():\n    print(f'batch labels: {lbs if TV_GE_0_16 else torch.sum(lbs, axis=0)}')\n    print(f'image size: {imgs[0].shape}')\n    for i in range(3):\n        ax = fig.add_subplot(3, 1, i + 1)\n        #print(np.rollaxis(imgs[i].numpy(), 0, 3).shape)\n        ax.imshow(np.rollaxis(imgs[i].numpy(), 0, 3))\n        ax.set_title(lbs[i])\n    break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:06:08.691568Z","iopub.execute_input":"2025-05-29T18:06:08.69185Z","iopub.status.idle":"2025-05-29T18:06:15.991465Z","shell.execute_reply.started":"2025-05-29T18:06:08.69182Z","shell.execute_reply":"2025-05-29T18:06:15.990458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q lion-pytorch adan-pytorch torch_optimizer","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:06:15.992596Z","iopub.execute_input":"2025-05-29T18:06:15.993028Z","iopub.status.idle":"2025-05-29T18:06:19.885174Z","shell.execute_reply.started":"2025-05-29T18:06:15.992993Z","shell.execute_reply":"2025-05-29T18:06:19.884329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install adan-pytorch lion-pytorch torch-optimizer","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:06:19.886228Z","iopub.execute_input":"2025-05-29T18:06:19.886574Z","iopub.status.idle":"2025-05-29T18:06:23.226081Z","shell.execute_reply.started":"2025-05-29T18:06:19.886548Z","shell.execute_reply":"2025-05-29T18:06:23.225167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install adan-pytorch lion-pytorch torch-optimizer","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T18:06:23.227167Z","iopub.execute_input":"2025-05-29T18:06:23.227506Z","iopub.status.idle":"2025-05-29T18:06:26.488999Z","shell.execute_reply.started":"2025-05-29T18:06:23.227479Z","shell.execute_reply":"2025-05-29T18:06:26.487893Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Modèle ResNet50ViTHybrid**","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport pytorch_lightning as pl\nimport timm\nfrom torchvision.models import resnet50, ResNet50_Weights\nimport torchmetrics\n\nclass CancerClassifierPL(pl.LightningModule):\n    def __init__(self, num_classes=6, class_weights=None, lr=5e-4):\n        super().__init__()\n        self.save_hyperparameters()\n\n        # === ResNet50 (local features) ===\n        resnet = resnet50(weights=ResNet50_Weights.DEFAULT)\n        self.resnet_features = nn.Sequential(*list(resnet.children())[:-2])  # Sans avgpool et fc\n        self.resnet_out_dim = 2048\n\n        # === ViT (global features) ===\n        self.vit = timm.create_model(\"vit_base_patch16_224\", pretrained=True)\n        self.vit.head = nn.Identity()\n        self.vit_out_dim = self.vit.embed_dim  # typiquement 768\n\n        # === Fusion & Classification ===\n        self.fusion_dim = self.resnet_out_dim + self.vit_out_dim\n        self.dropout = nn.Dropout(0.3)\n        self.classifier = nn.Sequential(\n            nn.Linear(self.fusion_dim, 512),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n\n        # === Loss ===\n        if class_weights is not None:\n            weights = torch.tensor(class_weights, dtype=torch.float32)\n            self.criterion = nn.CrossEntropyLoss(weight=weights, label_smoothing=0.1)\n        else:\n            self.criterion = nn.CrossEntropyLoss(label_smoothing=0.1)\n\n        # === Metrics ===\n        self.train_acc = torchmetrics.Accuracy(task=\"multiclass\", num_classes=num_classes)\n        self.val_acc = torchmetrics.Accuracy(task=\"multiclass\", num_classes=num_classes)\n        self.train_f1 = torchmetrics.F1Score(task=\"multiclass\", num_classes=num_classes)\n        self.val_f1 = torchmetrics.F1Score(task=\"multiclass\", num_classes=num_classes)\n\n    def forward(self, x):\n        # CNN branch (ResNet50)\n        cnn_features = self.resnet_features(x)  # (B, 2048, H', W')\n        cnn_pooled = torch.nn.functional.adaptive_avg_pool2d(cnn_features, (1, 1)).flatten(1)  # (B, 2048)\n\n        # ViT branch\n        vit_features = self.vit(x)  # (B, 768)\n\n        # Fusion\n        combined = torch.cat([cnn_pooled, vit_features], dim=1)  # (B, 2816)\n        output = self.classifier(self.dropout(combined))  # (B, num_classes)\n        return output\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.criterion(logits, y)\n        preds = torch.argmax(logits, dim=1)\n\n        self.train_acc(preds, y)\n        self.train_f1(preds, y)\n\n        self.log(\"train/loss\", loss, prog_bar=True)\n        self.log(\"train/acc\", self.train_acc, prog_bar=True)\n        self.log(\"train/f1\", self.train_f1, prog_bar=False)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.criterion(logits, y)\n        preds = torch.argmax(logits, dim=1)\n\n        self.val_acc(preds, y)\n        self.val_f1(preds, y)\n\n        self.log(\"val/loss\", loss, prog_bar=True)\n        self.log(\"val/acc\", self.val_acc, prog_bar=True)\n        self.log(\"val/f1\", self.val_f1, prog_bar=False)\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(self.parameters(), lr=self.hparams.lr)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)\n        return {\"optimizer\": optimizer, \"lr_scheduler\": scheduler}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T21:06:00.83797Z","iopub.execute_input":"2025-05-29T21:06:00.838361Z","iopub.status.idle":"2025-05-29T21:06:00.849827Z","shell.execute_reply.started":"2025-05-29T21:06:00.838298Z","shell.execute_reply":"2025-05-29T21:06:00.848987Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Entraînement avec PyTorch Lightning**","metadata":{}},{"cell_type":"code","source":"from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint\nfrom pytorch_lightning.loggers import CSVLogger\n\nmodel = CancerClassifierPL(num_classes=dm.num_classes)\nlogger = CSVLogger(\"logs\", name=\"cancer_subtype\")\n\n# Callback EarlyStopping\nearly_stopping = EarlyStopping(\n    monitor=\"val/acc\",   # métrique à surveiller\n    patience=5,           # nombre d'époques sans amélioration avant arrêt\n    mode=\"min\",           # \"min\" car on surveille une loss\n    verbose=True\n)\n\n# Callback ModelCheckpoint\ncheckpoint_callback = ModelCheckpoint(\n    monitor=\"val/acc\",           # même métrique que pour EarlyStopping\n    dirpath=\"checkpoints\",        # dossier où sauvegarder les checkpoints\n    filename=\"best-checkpoint\",   # nom du fichier checkpoint\n    save_top_k=1,                 # garder uniquement le meilleur\n    mode=\"min\",                   # \"min\" car on cherche à minimiser la loss\n    verbose=True\n)\n\ntrainer = pl.Trainer(\n    logger=logger,\n    callbacks=[early_stopping, checkpoint_callback],\n    max_epochs=50,\n    accelerator=\"gpu\", devices=1,\n    precision=\"16-mixed\",\n    log_every_n_steps=10\n)\n\ntrainer.fit(model, datamodule=dm)\n# --- Load best checkpoint ---\nbest_model_path = checkpoint_callback.best_model_path\nprint(f\"Loading best model from: {best_model_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T21:06:00.931788Z","iopub.execute_input":"2025-05-29T21:06:00.932002Z","iopub.status.idle":"2025-05-29T23:01:50.101294Z","shell.execute_reply.started":"2025-05-29T21:06:00.931981Z","shell.execute_reply":"2025-05-29T23:01:50.100121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\n\n# Charger les logs CSV (modifie ce chemin selon ton logger)\ncsv_path = f\"{logger.log_dir}/metrics.csv\"\ndf_logs = pd.read_csv(csv_path)\n\n# Moyenne des métriques par epoch\ndf_epoch = df_logs.groupby('epoch').mean()\n\n# --- Courbes de perte ---\ntrain_losses = df_epoch['train/loss'].tolist()\nval_losses = df_epoch['val/loss'].tolist()\nnum_epochs = len(train_losses)\n\nplt.figure(figsize=(12, 5))\n\nplt.subplot(1, 2, 1)  # 1 ligne, 2 colonnes, 1er plot\nplt.plot(range(1, num_epochs + 1), train_losses, label=\"Perte d'Entraînement\")\nplt.plot(range(1, num_epochs + 1), val_losses, label=\"Perte de Validation\")\nplt.xlabel(\"Époques\")\nplt.ylabel(\"Perte\")\nplt.title(\"Courbes de Perte d'Apprentissage et de Validation\")\nplt.legend()\nplt.grid(True)\n\n# Modification de la graduation de l'axe X (ticks)\nplt.xticks(ticks=range(1, num_epochs + 1), labels=[str(i) for i in range(1, num_epochs + 1)])\n\n# --- Courbes d'exactitude ---\ntrain_acc = df_epoch['train/acc'].tolist()\nval_acc = df_epoch['val/acc'].tolist()\n\ntrain_acc_with_zero = [0] + train_acc\nval_acc_with_zero = [0] + val_acc\nepochs_with_zero = range(0, num_epochs + 1)\n\nplt.subplot(1, 2, 2)  # 2e plot\nplt.plot(epochs_with_zero, train_acc_with_zero, label=\"Exactitude Entraînement\")\nplt.plot(epochs_with_zero, val_acc_with_zero, label=\"Exactitude Validation\")\n\n# Afficher max val accuracy\nmax_val_acc = max(val_acc)\nmax_val_epoch = val_acc.index(max_val_acc) + 1\n\nplt.scatter(max_val_epoch, max_val_acc, color='red')\nplt.text(max_val_epoch, max_val_acc + 0.01, f\"Max val acc: {max_val_acc:.3f}\", \n         ha='center', color='red')\n\nplt.xlabel(\"Époques\")\nplt.ylabel(\"Exactitude\")\nplt.title(\"Courbes d'Exactitude d'Apprentissage et de Validation\")\nplt.legend()\nplt.grid(True)\n\n# Ticks X avec 0 affiché mais numérotation qui commence à 1\nplt.xticks(ticks=epochs_with_zero, labels=[str(i) if i > 0 else \"0\" for i in epochs_with_zero])\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T23:01:50.104517Z","iopub.execute_input":"2025-05-29T23:01:50.10477Z","iopub.status.idle":"2025-05-29T23:01:50.644674Z","shell.execute_reply.started":"2025-05-29T23:01:50.104745Z","shell.execute_reply":"2025-05-29T23:01:50.64383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom sklearn.metrics import accuracy_score, f1_score, classification_report, confusion_matrix\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport numpy as np\n\n# S'assurer que le modèle est en mode évaluation et sur le bon device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\nmodel.eval()\n\n# Récupération du DataLoader de validation\nval_loader = dm.val_dataloader()\n\nall_preds = []\nall_targets = []\n\nwith torch.no_grad():\n    for x, y in val_loader:\n        x = x.to(device)\n        y = y.to(device)\n\n        logits = model(x)\n        preds = torch.argmax(logits, dim=1)\n\n        # Si les labels sont one-hot encodés\n        if dm.one_hot:\n            y = torch.argmax(y, dim=1)\n\n        all_preds.extend(preds.cpu().numpy())\n        all_targets.extend(y.cpu().numpy())\n\n# === Métriques ===\nacc = accuracy_score(all_targets, all_preds)\nf1_macro = f1_score(all_targets, all_preds, average='macro')\nf1_weighted = f1_score(all_targets, all_preds, average='weighted')\n\nprint(\"\\n=== Résultats ===\")\nprint(f\"Accuracy       : {acc:.4f}\")\nprint(f\"F1 Score Macro : {f1_macro:.4f}\")\nprint(f\"F1 Score Weighted : {f1_weighted:.4f}\")\n\nprint(\"\\n=== Rapport de classification ===\")\nprint(classification_report(all_targets, all_preds, digits=4))\n\n# === Matrice de confusion ===\ncm = confusion_matrix(all_targets, all_preds)\nplt.figure(figsize=(8, 6))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues')\nplt.title(\"Matrice de confusion\")\nplt.xlabel(\"Prédiction\")\nplt.ylabel(\"Vérité terrain\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T23:01:50.646098Z","iopub.execute_input":"2025-05-29T23:01:50.646357Z","iopub.status.idle":"2025-05-29T23:02:38.894165Z","shell.execute_reply.started":"2025-05-29T23:01:50.646334Z","shell.execute_reply":"2025-05-29T23:02:38.89316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torchvision import transforms\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport torch.nn.functional as F\n\ndef predict_image_with_display(img_path, checkpoint_path, class_names=None, image_size=224, num_classes=5, device=None):\n    device = device or ('cuda' if torch.cuda.is_available() else 'cpu')\n\n    # Désactive la protection anti DecompressionBomb\n    Image.MAX_IMAGE_PIXELS = None\n\n    # Charger et transformer l'image\n    image = Image.open(img_path).convert(\"RGB\")\n    transform = transforms.Compose([\n        transforms.Resize((image_size, image_size)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                             std=[0.229, 0.224, 0.225])\n    ])\n    img_tensor = transform(image).unsqueeze(0).to(device)\n\n    # Charger le modèle\n    model = CancerClassifierPL.load_from_checkpoint(checkpoint_path, num_classes=num_classes)\n    model.to(device)\n    model.eval()\n    model.freeze()\n\n    # Prédiction\n    with torch.no_grad():\n        logits = model(img_tensor)\n        probs = F.softmax(logits, dim=1).cpu().numpy().flatten()\n        pred_idx = probs.argmax()\n\n    pred_label = class_names[pred_idx] if class_names else str(pred_idx)\n\n    # Affichage\n    plt.figure(figsize=(8,4))\n    plt.subplot(1,2,1)\n    plt.imshow(image)\n    plt.axis('off')\n    plt.title(f\"Classe prédite : {pred_label}\", fontsize=14, color=\"green\")\n\n    plt.subplot(1,2,2)\n    y_pos = range(len(probs))\n    plt.barh(y_pos, probs)\n    plt.yticks(y_pos, class_names if class_names else [str(i) for i in y_pos])\n    plt.xlabel(\"Probabilité\")\n    plt.xlim(0,1)\n    plt.title(\"Distribution des probabilités\")\n\n    plt.tight_layout()\n    plt.show()\n\n    return pred_label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T23:02:38.895506Z","iopub.execute_input":"2025-05-29T23:02:38.895782Z","iopub.status.idle":"2025-05-29T23:02:38.90457Z","shell.execute_reply.started":"2025-05-29T23:02:38.895755Z","shell.execute_reply":"2025-05-29T23:02:38.903753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = ['LGSC', 'HGSC', 'CC', 'MC', 'EC' , 'Other']\nnum_classes = len(class_names)\n\nimg_path = \"/kaggle/input/UBC-OCEAN/test_images/41.png\"\nckpt_path = \"/kaggle/working/checkpoints/best-checkpoint-v1.ckpt\"\n\nprediction = predict_image_with_display(img_path, ckpt_path, class_names=class_names, num_classes=num_classes)\nprint(f\"✅ Résultat : {prediction}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T23:02:38.905539Z","iopub.execute_input":"2025-05-29T23:02:38.905868Z","iopub.status.idle":"2025-05-29T23:03:44.067793Z","shell.execute_reply.started":"2025-05-29T23:02:38.905844Z","shell.execute_reply":"2025-05-29T23:03:44.067052Z"}},"outputs":[],"execution_count":null}]}