{"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":"none","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":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"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},"outputs":[],"execution_count":null},{"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},"outputs":[],"execution_count":null},{"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},"outputs":[],"execution_count":null},{"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},"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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Prétraitement**","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},"outputs":[],"execution_count":null},{"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},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -U -q pytorch-lightning","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"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},"outputs":[],"execution_count":null},{"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},"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},"outputs":[],"execution_count":null},{"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},"outputs":[],"execution_count":null},{"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    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=None,\n        valid_transforms=None,\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_transforms = train_transforms\n        self.valid_transforms = valid_transforms\n        self.one_hot = one_hot\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\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            labels_lut=self.train_dataset.labels_lut)\n\n    def train_dataloader(self):\n        return DataLoader(self.train_dataset, batch_size=self.batch_size, num_workers=self.num_workers, shuffle=True)\n\n    def val_dataloader(self):\n        return DataLoader(self.valid_dataset, batch_size=self.batch_size, num_workers=self.num_workers, shuffle=False)\n\n    @property\n    def num_classes(self):\n        return len(self.train_dataset.labels_unique)\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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q lion-pytorch adan-pytorch torch_optimizer","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install adan-pytorch lion-pytorch torch-optimizer","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install adan-pytorch lion-pytorch torch-optimizer","metadata":{"trusted":true},"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 torchmetrics\nfrom torchvision.models import resnet50, ResNet50_Weights\n\nclass CancerClassifierPL(pl.LightningModule):\n    def __init__(self, num_classes=5, class_weights=None, lr=1e-3):\n        super().__init__()\n        self.save_hyperparameters()\n\n        self.cnn = resnet50(weights=ResNet50_Weights.DEFAULT)\n        self.cnn.fc = nn.Identity()\n\n        self.feature_dim = 2048\n        self.transformer_dim = 768\n        self.feature_adapter = nn.Linear(self.feature_dim, self.transformer_dim)\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=self.transformer_dim, nhead=8, dim_feedforward=2048,\n            dropout=0.1, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=6)\n        self.classifier = nn.Linear(self.transformer_dim, num_classes)\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)\n        else:\n            self.criterion = nn.CrossEntropyLoss()\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        features = self.cnn(x)\n        adapted = self.feature_adapter(features).unsqueeze(1)\n        encoded = self.transformer(adapted)\n        cls_token = encoded[:, 0, :]\n        return self.classifier(cls_token)\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)\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)\n\n    def predict_step(self, batch, batch_idx, dataloader_idx=0):\n        x = batch[0] if isinstance(batch, (list, tuple)) else batch\n        logits = self(x)\n        return torch.argmax(logits, dim=1)\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(self.parameters(), lr=self.hparams.lr)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer,\n            max_lr=self.hparams.lr,\n            steps_per_epoch=self.trainer.estimated_stepping_batches // self.trainer.max_epochs,\n            epochs=self.trainer.max_epochs,\n            pct_start=0.1,\n            anneal_strategy=\"cos\"\n        )\n        return {\"optimizer\": optimizer, \"lr_scheduler\": {\"scheduler\": scheduler, \"interval\": \"step\"}}\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pytorch_lightning.loggers import CSVLogger\nmodel = CancerClassifierPL(num_classes=dm.num_classes)\nlogger = CSVLogger(\"logs\", name=\"cancer_subtype\")\n\ntrainer = pl.Trainer(\n    logger=logger,\n    max_epochs=5,\n    accelerator=\"gpu\", devices=1,\n    precision=\"16-mixed\",\n    log_every_n_steps=10\n)\ntrainer.fit(model, datamodule=dm)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\n\nlog_path = \"logs/cancer_subtype/version_0/metrics.csv\"  # Adapte selon le nom/version\ndf = pd.read_csv(log_path)\n\nplt.figure(figsize=(12, 5))\n\n# Courbe de perte\nplt.subplot(1, 2, 1)\nplt.plot(df[\"step\"], df[\"train/loss\"], label=\"Train Loss\")\nplt.plot(df[\"step\"], df[\"val/loss\"], label=\"Val Loss\")\nplt.xlabel(\"Step\")\nplt.ylabel(\"Loss\")\nplt.title(\"Courbe de perte\")\nplt.legend()\n\n# Courbe d'accuracy\nplt.subplot(1, 2, 2)\nplt.plot(df[\"step\"], df[\"train/acc\"], label=\"Train Accuracy\")\nplt.plot(df[\"step\"], df[\"val/acc\"], label=\"Val Accuracy\")\nplt.xlabel(\"Step\")\nplt.ylabel(\"Accuracy\")\nplt.title(\"Courbe de précision\")\nplt.legend()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nfrom torchmetrics.functional import accuracy, f1_score\nfrom sklearn.metrics import classification_report, confusion_matrix\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n# Charger le modèle\nmodel = CancerClassifierPL.load_from_checkpoint(\"path/to/best_model.ckpt\")\n\n# Charger les données de test\ntest_loader = dm.val_dataloader()  # Ou un vrai test_dataloader si tu en as\n\nall_preds = []\nall_labels = []\n\nmodel.eval()\nmodel.to(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nwith torch.no_grad():\n    for batch in test_loader:\n        x, y = batch\n        x, y = x.to(model.device), y.to(model.device)\n        preds = model(x)\n        preds = torch.argmax(preds, dim=1)\n\n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(y.cpu().numpy())\n\n# Classification report\nprint(\"Classification Report:\")\nprint(classification_report(all_labels, all_preds))\n\n# Confusion matrix\ncm = confusion_matrix(all_labels, all_preds)\nplt.figure(figsize=(6, 6))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues')\nplt.xlabel(\"Prédictions\")\nplt.ylabel(\"Vrais labels\")\nplt.title(\"Matrice de confusion\")\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom sklearn.metrics import classification_report, confusion_matrix, ConfusionMatrixDisplay\nimport matplotlib.pyplot as plt\n\ndef evaluate_model(model, dataloader):\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n    model.eval()\n\n    all_preds = []\n    all_targets = []\n\n    with torch.no_grad():\n        for batch in dataloader:\n            x, y = batch\n            x, y = x.to(device), y.to(device)\n            logits = model(x)\n            preds = torch.argmax(logits, dim=1)\n            all_preds.extend(preds.cpu().numpy())\n            all_targets.extend(y.cpu().numpy())\n\n    # 🧾 Rapport de classification\n    print(\"\\n📊 Rapport de classification :\")\n    print(classification_report(all_targets, all_preds, target_names=class_names))\n\n    # 📉 Matrice de confusion\n    cm = confusion_matrix(all_targets, all_preds)\n    disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=class_names)\n\n    plt.figure(figsize=(8, 6))\n    disp.plot(cmap=plt.cm.Blues, xticks_rotation=45, values_format='d')\n    plt.title(\"Matrice de confusion\")\n    plt.grid(False)\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Charger le meilleur modèle sauvegardé\nmodel = ResNet50ViTHybridPL(num_classes=5, lr=1e-4)\n\n# Évaluer sur les données de test\nevaluate_model(model, dm.val_dataloader())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Cross validation ","metadata":{}},{"cell_type":"markdown","source":"Prediction","metadata":{}},{"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = ['LGSC', 'HGSC', 'CC', 'MC', 'EC']\nnum_classes = len(class_names)\n\nimg_path = \"/kaggle/input/UBC-OCEAN/test_images/41.png\"\nckpt_path = \"/kaggle/working/logs/cancer_classifier/version_2/checkpoints/best_model.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},"outputs":[],"execution_count":null}]}