{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":3486,"databundleVersionId":31310,"sourceType":"competition"},{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"},{"sourceId":1834160,"sourceType":"datasetVersion","datasetId":333968}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport 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        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:44:33.071451Z","iopub.execute_input":"2025-05-24T16:44:33.071998Z","iopub.status.idle":"2025-05-24T16:44:44.737322Z","shell.execute_reply.started":"2025-05-24T16:44:33.071974Z","shell.execute_reply":"2025-05-24T16:44:44.736392Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -U albumentations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:09:56.263768Z","iopub.execute_input":"2025-05-24T16:09:56.264192Z","iopub.status.idle":"2025-05-24T16:10:01.534021Z","shell.execute_reply.started":"2025-05-24T16:09:56.264161Z","shell.execute_reply":"2025-05-24T16:10:01.533265Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!nvidia-smi","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:48:06.990757Z","iopub.execute_input":"2025-05-24T16:48:06.991068Z","iopub.status.idle":"2025-05-24T16:48:07.637602Z","shell.execute_reply.started":"2025-05-24T16:48:06.991047Z","shell.execute_reply":"2025-05-24T16:48:07.636574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch import device\nfrom torch.cuda import is_available\nimport torch.nn as nn\nimport os\n\n\nDATASET_PATH = \"/kaggle/input/cassava-leaf-disease-classification/\"\nTRAIN_PATH = \"/kaggle/input/cassava-leaf-disease-classification/train_images\" \nTEST_PATH = \"/kaggle/input/cassava-leaf-disease-classification/test_images\" \nLABELS_JSON = \"/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json\"\n\nNUM_EPOCHS = 5\nTARGET_SIZE = 224\nBATCH_SIZE = 64\nRANDOM_STATE = 42\nNUM_WORKERS = min(4, os.cpu_count())\nDEVICE = device(\"cuda\" if is_available() else \"cpu\")\nCRITERION = nn.CrossEntropyLoss().to(DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:21:03.409119Z","iopub.execute_input":"2025-05-24T19:21:03.409404Z","iopub.status.idle":"2025-05-24T19:21:03.414568Z","shell.execute_reply.started":"2025-05-24T19:21:03.409384Z","shell.execute_reply":"2025-05-24T19:21:03.41393Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport numpy as np\nimport torch\n\n\ndef set_seed(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed) \n\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False \n\nset_seed(seed=RANDOM_STATE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:10:04.78571Z","iopub.execute_input":"2025-05-24T16:10:04.786103Z","iopub.status.idle":"2025-05-24T16:10:04.795882Z","shell.execute_reply.started":"2025-05-24T16:10:04.786081Z","shell.execute_reply":"2025-05-24T16:10:04.795172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import csv\nimport os\nfrom PIL import Image\nfrom itertools import islice\n\ntrain_dataset = []\n\nwith open(os.path.join(DATASET_PATH, \"train.csv\"), \"r\") as file_obj:\n    next(file_obj)\n    reader = csv.reader(file_obj)\n    \n    for row in islice(reader, 10_000):\n        image_path = os.path.join(TRAIN_PATH, row[0])\n        if os.path.exists(image_path):\n            image = Image.open(image_path).convert(\"RGB\")\n            train_dataset.append([np.array(image), int(row[1])])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:10:06.419594Z","iopub.execute_input":"2025-05-24T16:10:06.420276Z","iopub.status.idle":"2025-05-24T16:11:53.16352Z","shell.execute_reply.started":"2025-05-24T16:10:06.420251Z","shell.execute_reply":"2025-05-24T16:11:53.16284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nfrom pprint import pprint\n\n\nwith open (LABELS_JSON, \"r\") as json_file: \n    id2label = json.load(json_file)\n    id2label = {key: value.split()[-1].strip(\"()\") for key, value in id2label.items()}\n    label2id = {value: key for key, value in id2label.items()}\n\npprint(id2label)\nprint()\npprint(label2id)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:53.164849Z","iopub.execute_input":"2025-05-24T16:11:53.165083Z","iopub.status.idle":"2025-05-24T16:11:53.172108Z","shell.execute_reply.started":"2025-05-24T16:11:53.165061Z","shell.execute_reply":"2025-05-24T16:11:53.171541Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import List\nimport matplotlib.pyplot as plt\n\n\ndef show_files(indices: List[int], images: List[list[np.array, int]]) -> None:\n    n = len(indices)\n    images = [images[i] for i in indices]\n    n_col = 3\n    n_rows = (n + n_col - 1) // n_col  \n    plt.figure(figsize=(15, 5 * n_rows))\n    for i in range(min(n, len(images))):\n        img, label = images[i]\n        plt.subplot(n_rows, n_col, i + 1)\n        plt.imshow(img)\n        plt.title(f\"Label: {id2label[str(label)]}\")\n        plt.axis('off')\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:53.173022Z","iopub.execute_input":"2025-05-24T16:11:53.173276Z","iopub.status.idle":"2025-05-24T16:11:53.183509Z","shell.execute_reply.started":"2025-05-24T16:11:53.173248Z","shell.execute_reply":"2025-05-24T16:11:53.182782Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"indices = np.random.choice(range(len(train_dataset)), 6)\nshow_files(\n    indices=indices,\n    images=train_dataset\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:53.185483Z","iopub.execute_input":"2025-05-24T16:11:53.185722Z","iopub.status.idle":"2025-05-24T16:11:54.460224Z","shell.execute_reply.started":"2025-05-24T16:11:53.185705Z","shell.execute_reply":"2025-05-24T16:11:54.459052Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Dataset\nfrom typing import Tuple\n\n\nclass FlowersDataset(Dataset):\n    def __init__(self, images_list: List[list[np.array, int]], transform = None):\n        self._images_list = images_list\n        self._transform = transform\n    \n    def __len__(self) -> int:\n        return len(self._images_list)\n\n    def __getitem__(self, index) -> Tuple[np.array, np.array]:\n        image, label = self._images_list[index]\n\n        if self._transform:\n            transformed_data = self._transform(image=image)\n            image = transformed_data['image']\n        \n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:54.461496Z","iopub.execute_input":"2025-05-24T16:11:54.462235Z","iopub.status.idle":"2025-05-24T16:11:54.470383Z","shell.execute_reply.started":"2025-05-24T16:11:54.462188Z","shell.execute_reply":"2025-05-24T16:11:54.469743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A\n\n\ntrain_transform = A.Compose([\n    A.SmallestMaxSize(max_size=TARGET_SIZE, p=1.0),\n    A.RandomCrop(height=TARGET_SIZE, width=TARGET_SIZE, p=1.0),\n    A.HorizontalFlip(p=0.5),\n    A.Blur(p=0.3),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), max_pixel_value=255.0, p=1.0),\n    A.ToTensorV2()\n])\n\nval_transform = A.Compose([\n    A.SmallestMaxSize(max_size=TARGET_SIZE, p=1.0),\n    A.CenterCrop(height=TARGET_SIZE, width=TARGET_SIZE, p=1.0),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    A.ToTensorV2(),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:54.471233Z","iopub.execute_input":"2025-05-24T16:11:54.471491Z","iopub.status.idle":"2025-05-24T16:11:56.019751Z","shell.execute_reply.started":"2025-05-24T16:11:54.471471Z","shell.execute_reply":"2025-05-24T16:11:56.019178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n\nX_train, X_temp = train_test_split(\n    train_dataset,\n    test_size=0.2,\n    random_state=RANDOM_STATE,\n)\n\nX_test, X_valid = train_test_split(\n    X_temp,\n    test_size=0.7,\n    random_state=RANDOM_STATE,\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:56.02072Z","iopub.execute_input":"2025-05-24T16:11:56.021016Z","iopub.status.idle":"2025-05-24T16:11:56.943725Z","shell.execute_reply.started":"2025-05-24T16:11:56.020999Z","shell.execute_reply":"2025-05-24T16:11:56.942953Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Размерность обучающей выборки: {}\".format(len(X_train)))\nprint(\"Размерность валидацинной выборки: {}\".format(len(X_valid)))\nprint(\"Размерность тестовой выборки: {}\".format(len(X_test)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:56.944669Z","iopub.execute_input":"2025-05-24T16:11:56.945358Z","iopub.status.idle":"2025-05-24T16:11:56.949606Z","shell.execute_reply.started":"2025-05-24T16:11:56.945331Z","shell.execute_reply":"2025-05-24T16:11:56.948899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Без aug\ntrain_dataset_without_aug = FlowersDataset(\n    images_list=X_train,\n    transform=val_transform\n)\n\n# С aug\ntrain_dataset_with_aug = FlowersDataset(\n    images_list=X_train,\n    transform=train_transform\n)\n\n# Валидационная и Тестовая выборки\nvalid_dataset = FlowersDataset(\n    images_list=X_valid,\n    transform=val_transform\n)\n\ntest_dataset = FlowersDataset(\n    images_list=X_test,\n    transform=val_transform\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:56.950355Z","iopub.execute_input":"2025-05-24T16:11:56.950593Z","iopub.status.idle":"2025-05-24T16:11:56.961407Z","shell.execute_reply.started":"2025-05-24T16:11:56.950577Z","shell.execute_reply":"2025-05-24T16:11:56.960838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset_without_aug[0][0].size()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:56.963436Z","iopub.execute_input":"2025-05-24T16:11:56.963662Z","iopub.status.idle":"2025-05-24T16:11:56.979644Z","shell.execute_reply.started":"2025-05-24T16:11:56.963646Z","shell.execute_reply":"2025-05-24T16:11:56.978902Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader_without_aug = torch.utils.data.DataLoader(\n    dataset=train_dataset_without_aug,\n    batch_size=BATCH_SIZE,\n    num_workers=NUM_WORKERS,\n    shuffle=True\n)\ntrain_loader_with_aug = torch.utils.data.DataLoader(\n    dataset=train_dataset_with_aug,\n    batch_size=BATCH_SIZE,\n    num_workers=NUM_WORKERS,\n    shuffle=True\n)\n\nvalid_loader = torch.utils.data.DataLoader(\n    dataset=valid_dataset,\n    batch_size=1,\n    num_workers=NUM_WORKERS\n)\n\ntest_lodaer = torch.utils.data.DataLoader(\n    dataset=test_dataset,\n    batch_size=1,\n    num_workers=NUM_WORKERS\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:11:56.980287Z","iopub.execute_input":"2025-05-24T16:11:56.980524Z","iopub.status.idle":"2025-05-24T16:11:56.988017Z","shell.execute_reply.started":"2025-05-24T16:11:56.980508Z","shell.execute_reply":"2025-05-24T16:11:56.987222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\n\nclass CNNModel(nn.Module):\n    def __init__(self,\n                 n_channels: int,\n                 id2label: dict,\n                 is_batch_norm: bool = True,\n                 is_dropout: bool = True):\n        super(CNNModel, self).__init__()\n        \n        channels = [3, n_channels, n_channels * 2, n_channels * 4]\n        conv_blocks = []\n        for in_ch, out_ch in zip(channels[:-1], channels[1:]):\n            layers = [nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=1, padding=1)]\n            if is_batch_norm:\n                layers.append(nn.BatchNorm2d(out_ch))\n            layers.append(nn.ReLU())             \n            if is_dropout:\n                layers.append(nn.Dropout2d(p=0.3))  \n            layers.append(nn.MaxPool2d(kernel_size=2, stride=2))\n            conv_blocks.append(nn.Sequential(*layers))\n        \n        self.features = nn.Sequential(*conv_blocks)\n        \n        flattened_size = channels[-1] * 28 * 28\n        \n        self.fc1 = nn.Linear(flattened_size, 256)\n        self.fc2 = nn.Linear(256, 64)\n        self.fc3 = nn.Linear(64, len(id2label))\n        self.is_dropout = is_dropout\n\n    def forward(self, x):\n        x = self.features(x)\n        x = x.view(x.size(0), -1)\n        x = F.relu(self.fc1(x))         \n        if self.is_dropout:\n            x = F.dropout(x, p=0.5)      \n        x = F.relu(self.fc2(x))\n        if self.is_dropout:\n            x = F.dropout(x, p=0.5)\n        x = self.fc3(x)\n        return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:29:14.288235Z","iopub.execute_input":"2025-05-24T16:29:14.28888Z","iopub.status.idle":"2025-05-24T16:29:14.297454Z","shell.execute_reply.started":"2025-05-24T16:29:14.288847Z","shell.execute_reply":"2025-05-24T16:29:14.296673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cnn_model = CNNModel(\n    n_channels=32,\n    id2label=id2label,\n    is_batch_norm=False,\n    is_dropout=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:29:18.133068Z","iopub.execute_input":"2025-05-24T16:29:18.133823Z","iopub.status.idle":"2025-05-24T16:29:18.355688Z","shell.execute_reply.started":"2025-05-24T16:29:18.133795Z","shell.execute_reply":"2025-05-24T16:29:18.355085Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cnn_model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T16:29:18.568084Z","iopub.execute_input":"2025-05-24T16:29:18.56831Z","iopub.status.idle":"2025-05-24T16:29:18.573812Z","shell.execute_reply.started":"2025-05-24T16:29:18.568291Z","shell.execute_reply":"2025-05-24T16:29:18.573134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torchvision.models import mobilenet_v2, MobileNet_V2_Weights\n\n\nclass TrainModels:\n    def __init__(self,\n                 exp,\n                 num_epochs,\n                 train_loader,\n                 train_loader_aug,\n                 valid_loader,\n                 test_loader,\n                 CNNModel,\n                 id2label,\n                 device,\n                 criterion):\n        self.exp = exp\n        self.device = device\n        self.criterion = criterion\n        self.num_epochs = num_epochs\n        \n        self.train_loader = train_loader_aug if exp['aug'] else train_loader\n        self.valid_loader = valid_loader\n        self.test_loader = test_loader\n\n        print(f\"Модель: {exp['name']}\")\n        if exp.get('pretrained', False):\n            self.model = mobilenet_v2(weights=MobileNet_V2_Weights.IMAGENET1K_V2)\n            in_feat = self.model.classifier[1].in_features\n            self.model.classifier[1] = nn.Linear(in_feat, len(id2label))\n        else:\n            self.model = CNNModel(\n                n_channels=32,\n                id2label=id2label,\n                is_batch_norm=True,\n                is_dropout=True\n            )\n        self.model = self.model.to(self.device)\n\n        if exp['optim'] == 'SGD':\n            self.optimizer = optim.SGD(self.model.parameters(), lr=1e-2, momentum=0.9)\n        else:\n            self.optimizer = optim.Adam(self.model.parameters(), lr=1e-3)\n            \n        self.scheduler = (lr_scheduler.StepLR(self.optimizer, step_size=5, gamma=0.1)\n                          if exp['sched'] else None)\n\n    def _train_step(self, x, y):\n        x, y = x.to(self.device), y.to(self.device)\n        self.optimizer.zero_grad()\n        out = self.model(x)\n        loss = self.criterion(out, y)\n        loss.backward()\n        self.optimizer.step()\n        preds = out.argmax(dim=1)\n        return loss.item(), preds.eq(y).sum().item(), x.size(0)\n\n    def _eval_step(self, x, y):\n        x, y = x.to(self.device), y.to(self.device)\n        out = self.model(x)\n        loss = self.criterion(out, y)\n        preds = out.argmax(dim=1)\n        return loss.item(), preds.eq(y).sum().item(), x.size(0)\n\n    def train_epoch(self):\n        self.model.train()\n        total_loss, total_correct, total_samples = 0, 0, 0\n        for x, y in self.train_loader:\n            loss, correct, n = self._train_step(x, y)\n            total_loss += loss * n\n            total_correct += correct\n            total_samples += n\n        return total_loss / total_samples, total_correct / total_samples\n\n    def eval_epoch(self, loader):\n        self.model.eval()\n        total_loss, total_correct, total_samples = 0, 0, 0\n        with torch.no_grad():\n            for x, y in loader:\n                loss, correct, n = self._eval_step(x, y)\n                total_loss += loss * n\n                total_correct += correct\n                total_samples += n\n        return total_loss / total_samples, total_correct / total_samples\n\n    def run(self):\n        history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': []}\n        for epoch in range(1, self.num_epochs + 1):\n            tl, ta = self.train_epoch()\n            vl, va = self.eval_epoch(self.valid_loader)\n\n            history['train_loss'].append(tl)\n            history['train_acc'].append(ta)\n            history['val_loss'].append(vl)\n            history['val_acc'].append(va)\n\n            print(f\"Epoch {epoch}/{self.num_epochs}: \"\n                  f\"train_loss={tl:.4f}, train_acc={ta:.4f}, \"\n                  f\"val_loss={vl:.4f}, val_acc={va:.4f}\")\n\n            if self.scheduler:\n                self.scheduler.step()\n\n        test_loss, test_acc = self.eval_epoch(self.test_loader)\n        print(f\"Результат на тестовой выборке: loss={test_loss:.4f}, acc={test_acc:.4f}\")\n        return history, test_acc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:24:21.062898Z","iopub.execute_input":"2025-05-24T19:24:21.063694Z","iopub.status.idle":"2025-05-24T19:24:21.078336Z","shell.execute_reply.started":"2025-05-24T19:24:21.063666Z","shell.execute_reply":"2025-05-24T19:24:21.077571Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"exp_config = [\n    {\n        \"name\":\"Custom + Adam+RLP\",  \n        \"optim\":\"Adam\",   \n        \"sched\":True,  \n        \"aug\":False\n    },\n    \n    {\n        \"name\":\"MobileNet V2 + Adam\",   \n        \"optim\":\"Adam\",   \n        \"sched\":False, \n        \"aug\":False, \n        \"pretrained\":True\n    },\n    \n    {\n        \"name\":\"Custom + SGD\", \n        \"optim\":\"SGD\",    \n        \"sched\":False, \n        \"aug\":False\n    },\n\n    {   \"name\":\"Custom + Adam\", \n        \"optim\":\"Adam\", \n        \"sched\":False, \n        \"aug\":False\n    },\n\n    {\n        \"name\":\"Custom + SGD+Aug\",  \n        \"optim\":\"SGD\",    \n        \"sched\":False, \n        \"aug\":True\n    }\n]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:24:22.925077Z","iopub.execute_input":"2025-05-24T19:24:22.925363Z","iopub.status.idle":"2025-05-24T19:24:22.930279Z","shell.execute_reply.started":"2025-05-24T19:24:22.925344Z","shell.execute_reply":"2025-05-24T19:24:22.929541Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"histories = {}\nresults = []\n\nfor exp in exp_config:\n    trainer = TrainModels(\n        exp=exp,\n        num_epochs=NUM_EPOCHS,\n        train_loader=train_loader_without_aug,\n        train_loader_aug=train_loader_with_aug,\n        valid_loader=valid_loader,\n        test_loader=test_lodaer,\n        CNNModel=CNNModel,\n        id2label=id2label,\n        device=DEVICE,\n        criterion=CRITERION\n    )\n    history, test_acc = trainer.run()\n    histories[exp['name']] = history\n    results.append({\n        'experiment': exp['name'],\n        'train_acc': history['train_acc'][-1],\n        'test_acc': test_acc\n    })\n    print()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:24:23.763278Z","iopub.execute_input":"2025-05-24T19:24:23.763546Z","iopub.status.idle":"2025-05-24T19:38:50.730385Z","shell.execute_reply.started":"2025-05-24T19:24:23.763529Z","shell.execute_reply":"2025-05-24T19:38:50.729617Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n\nexp_result = pd.DataFrame(results).set_index('experiment')\nexp_result.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:43:37.299549Z","iopub.execute_input":"2025-05-24T19:43:37.30029Z","iopub.status.idle":"2025-05-24T19:43:37.314639Z","shell.execute_reply.started":"2025-05-24T19:43:37.300261Z","shell.execute_reply":"2025-05-24T19:43:37.313996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best = max(results, key=lambda x: x['test_acc'])['experiment']\nbest_hist = histories[best]\n\nplt.figure(figsize=(8,5))\nplt.plot(best_hist['train_loss'], label='train loss')\nplt.plot(best_hist['val_loss'], label='val loss')\nplt.title(f\"Лучший результат: {best}\")\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:43:40.101358Z","iopub.execute_input":"2025-05-24T19:43:40.10167Z","iopub.status.idle":"2025-05-24T19:43:40.29732Z","shell.execute_reply.started":"2025-05-24T19:43:40.101647Z","shell.execute_reply":"2025-05-24T19:43:40.296558Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"activations = {}\n\ndef get_activation(name):\n    def hook(model, input, output):\n        activations[name] = output.detach()\n    return hook\n\nmodel.layer1.register_forward_hook(get_activation('layer1'))\nmodel.layer2.register_forward_hook(get_activation('layer2'))\n\nmodel.eval()\nwith torch.no_grad():\n    x = sample_input.unsqueeze(0).to(device)   \n    _ = model(x)\n\nact = activations['layer1']   \n\nn_channels = min(8, act.shape[1])\nfig, axes = plt.subplots(1, n_channels, figsize=(n_channels*2, 2))\nfor i in range(n_channels):\n    axes[i].imshow(act[0, i].cpu(), cmap='viridis')\n    axes[i].axis('off')\nplt.suptitle(\"Активации layer1 (первые каналы)\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Задание 1.2 CNN локализация лицевых точек","metadata":{}},{"cell_type":"markdown","source":"https://www.kaggle.com/competitions/facial-keypoints-detection/overview","metadata":{}},{"cell_type":"code","source":"import cv2\nfrom torch.optim import lr_scheduler\nfrom torchvision import transforms\nfrom torchvision.models import resnet18, ResNet18_Weights\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm.notebook import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:48:37.691641Z","iopub.execute_input":"2025-05-24T19:48:37.692036Z","iopub.status.idle":"2025-05-24T19:48:37.696643Z","shell.execute_reply.started":"2025-05-24T19:48:37.692001Z","shell.execute_reply":"2025-05-24T19:48:37.695928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_EPOCHS = 15 \nTARGET_SIZE = 96\nBATCH_SIZE = 64\n\nset_seed(seed=RANDOM_STATE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:52:54.977068Z","iopub.execute_input":"2025-05-24T19:52:54.977641Z","iopub.status.idle":"2025-05-24T19:52:54.983134Z","shell.execute_reply.started":"2025-05-24T19:52:54.977613Z","shell.execute_reply":"2025-05-24T19:52:54.982367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = '/kaggle/input/facial-keypoints-detection/'\ntrain_df_full = pd.read_csv(os.path.join(DATA_DIR, 'training.zip'), compression='zip')\ntest_df_raw = pd.read_csv(os.path.join(DATA_DIR, 'test.zip'), compression='zip')\n\nkeypoint_cols = train_df_full.columns[:-1].tolist()\nNUM_KEYPOINTS_COORDS = len(keypoint_cols) # 30\nNAN_PLACEHOLDER = -1.0 ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:52:55.245996Z","iopub.execute_input":"2025-05-24T19:52:55.24621Z","iopub.status.idle":"2025-05-24T19:53:00.238732Z","shell.execute_reply.started":"2025-05-24T19:52:55.246194Z","shell.execute_reply":"2025-05-24T19:53:00.238149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MaskedMSELoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.criterion = nn.MSELoss(reduction='none') \n\n    def forward(self, preds, targets, valid_mask):\n        loss_per_coord = self.criterion(preds, targets)\n        masked_loss_per_coord = loss_per_coord * valid_mask\n\n        num_valid_coords = valid_mask.sum()\n        if num_valid_coords > 0:\n            mean_loss = masked_loss_per_coord.sum() / num_valid_coords\n        else:\n            mean_loss = torch.tensor(0.0, device=preds.device, dtype=preds.dtype)\n        return mean_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:53:00.240118Z","iopub.execute_input":"2025-05-24T19:53:00.240371Z","iopub.status.idle":"2025-05-24T19:53:00.24566Z","shell.execute_reply.started":"2025-05-24T19:53:00.240353Z","shell.execute_reply":"2025-05-24T19:53:00.244963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CRITERION = MaskedMSELoss().to(DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:53:00.246478Z","iopub.execute_input":"2025-05-24T19:53:00.24671Z","iopub.status.idle":"2025-05-24T19:53:00.258059Z","shell.execute_reply.started":"2025-05-24T19:53:00.246685Z","shell.execute_reply":"2025-05-24T19:53:00.257342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FacialKeypointsDataset(Dataset):\n    def __init__(self, dataframe, keypoint_cols, transform=None, nan_placeholder=NAN_PLACEHOLDER):\n        self.dataframe = dataframe\n        self.keypoint_cols = keypoint_cols\n        self.transform = transform\n        self.nan_placeholder = nan_placeholder\n\n        self.images_str = self.dataframe['Image'].values \n        self.keypoints_data = self.dataframe[self.keypoint_cols].values.astype(np.float32) \n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        img_str = self.images_str[idx]\n        image = np.array(img_str.split(), dtype=np.uint8).reshape(TARGET_SIZE, TARGET_SIZE)\n        \n        keypoints_original = self.keypoints_data[idx].copy()\n\n        valid_mask_flat = ~np.isnan(keypoints_original)\n        valid_mask_flat = valid_mask_flat.astype(np.float32)\n\n        keypoints_filled = np.nan_to_num(keypoints_original, nan=self.nan_placeholder)\n\n        keypoints_for_aug = []\n        for i in range(0, len(keypoints_filled), 2):\n            keypoints_for_aug.append([keypoints_filled[i], keypoints_filled[i+1]])\n        \n        image_rgb = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)\n\n        processed_keypoints_flat = keypoints_filled\n\n        if self.transform:\n            transformed = self.transform(image=image_rgb, keypoints=keypoints_for_aug)\n            image_rgb = transformed['image']\n            \n            transformed_keypoints_list = transformed['keypoints']\n            \n            temp_processed_keypoints = np.full(self.keypoints_data.shape[1], self.nan_placeholder, dtype=np.float32)\n            \n            idx_kp = 0\n            for i in range(0, len(temp_processed_keypoints), 2):\n                if idx_kp < len(transformed_keypoints_list):\n                    temp_processed_keypoints[i] = transformed_keypoints_list[idx_kp][0]\n                    temp_processed_keypoints[i+1] = transformed_keypoints_list[idx_kp][1]\n                    idx_kp += 1\n\n            processed_keypoints_flat = temp_processed_keypoints\n            processed_keypoints_flat = (processed_keypoints_flat / float(TARGET_SIZE)) * 2.0 - 1.0\n        else: \n            image_rgb = torch.from_numpy(image_rgb.transpose(2,0,1)).float() / 255.0 \n            processed_keypoints_flat = (processed_keypoints_flat / float(TARGET_SIZE)) * 2.0 - 1.0\n\n\n        return image_rgb, torch.tensor(processed_keypoints_flat, dtype=torch.float32), torch.tensor(valid_mask_flat, dtype=torch.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:53:00.259546Z","iopub.execute_input":"2025-05-24T19:53:00.259833Z","iopub.status.idle":"2025-05-24T19:53:00.276482Z","shell.execute_reply.started":"2025-05-24T19:53:00.259817Z","shell.execute_reply":"2025-05-24T19:53:00.275753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"kp_params = A.KeypointParams(format='xy', remove_invisible=False)\n\ntrain_transform = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.Rotate(limit=25, p=0.6, border_mode=cv2.BORDER_CONSTANT, value=0),\n    A.RandomBrightnessContrast(brightness_limit=0.25, contrast_limit=0.25, p=0.6),\n    A.GaussNoise(var_limit=(10.0, 60.0), p=0.4),\n    A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.1, rotate_limit=0, p=0.5, border_mode=cv2.BORDER_CONSTANT, value=0),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), max_pixel_value=255.0),\n    ToTensorV2()\n], keypoint_params=kp_params)\n\nval_transform = A.Compose([\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), max_pixel_value=255.0),\n    ToTensorV2(),\n], keypoint_params=kp_params)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:53:00.2772Z","iopub.execute_input":"2025-05-24T19:53:00.277398Z","iopub.status.idle":"2025-05-24T19:53:00.295776Z","shell.execute_reply.started":"2025-05-24T19:53:00.277379Z","shell.execute_reply":"2025-05-24T19:53:00.295122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_sub_df, val_sub_df = train_test_split(train_df_full, test_size=0.2, random_state=RANDOM_STATE)\n\ntrain_dataset_no_aug = FacialKeypointsDataset(train_sub_df, keypoint_cols, transform=val_transform)\ntrain_dataset_with_aug = FacialKeypointsDataset(train_sub_df, keypoint_cols, transform=train_transform)\nval_dataset = FacialKeypointsDataset(val_sub_df, keypoint_cols, transform=val_transform)\n\ntrain_loader_without_aug = torch.utils.data.DataLoader(train_dataset_no_aug, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)\ntrain_loader_with_aug = torch.utils.data.DataLoader(train_dataset_with_aug, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)\nvalid_loader = torch.utils.data.DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:53:00.296503Z","iopub.execute_input":"2025-05-24T19:53:00.296751Z","iopub.status.idle":"2025-05-24T19:53:00.309739Z","shell.execute_reply.started":"2025-05-24T19:53:00.29673Z","shell.execute_reply":"2025-05-24T19:53:00.309111Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CustomCNNKeypoints(nn.Module):\n    def __init__(self, num_keypoints_coords, n_channels_start=32, is_batch_norm=True, is_dropout=True):\n        super().__init__()\n        self.num_keypoints_coords = num_keypoints_coords \n\n        channels = [3, n_channels_start, n_channels_start * 2, n_channels_start * 4, n_channels_start * 8]\n        conv_blocks = []\n        current_size = TARGET_SIZE \n\n        for i in range(len(channels) - 1):\n            in_ch = channels[i]\n            out_ch = channels[i+1]\n            layers = [nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=1, padding=1)]\n            if is_batch_norm:\n                layers.append(nn.BatchNorm2d(out_ch))\n            layers.append(nn.ReLU(inplace=True))\n            if is_dropout and i < len(channels) - 2 : \n                 layers.append(nn.Dropout2d(p=0.25)) \n            layers.append(nn.MaxPool2d(kernel_size=2, stride=2))\n            conv_blocks.append(nn.Sequential(*layers))\n            current_size = current_size // 2 \n        self.features = nn.Sequential(*conv_blocks)\n        \n        flattened_size = channels[-1] * current_size * current_size\n\n        self.fc_layers = nn.Sequential(\n            nn.Linear(flattened_size, 1024),\n            nn.ReLU(inplace=True),\n            nn.Dropout(p=0.5 if is_dropout else 0.0), \n            nn.Linear(1024, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(p=0.5 if is_dropout else 0.0),\n            nn.Linear(512, self.num_keypoints_coords),\n            nn.Tanh() \n        )\n\n    def forward(self, x):\n        x = self.features(x)\n        x = x.view(x.size(0), -1) \n        x = self.fc_layers(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:53:00.31058Z","iopub.execute_input":"2025-05-24T19:53:00.310782Z","iopub.status.idle":"2025-05-24T19:53:00.318827Z","shell.execute_reply.started":"2025-05-24T19:53:00.310766Z","shell.execute_reply":"2025-05-24T19:53:00.318227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TrainKeypointModels:\n    def __init__(self, experiment_config, num_epochs, train_loader_no_aug, train_loader_with_aug,\n                 valid_loader, CustomCNNKeypoints, num_keypoints_coords, device, criterion):\n        self.exp_config = experiment_config\n        self.device = device\n        self.criterion = criterion\n        self.num_epochs = num_epochs\n        self.num_keypoints_coords = num_keypoints_coords\n        \n        self.train_loader = train_loader_with_aug if self.exp_config.get('aug', False) else train_loader_no_aug\n        self.valid_loader = valid_loader\n\n        print(f\"Эксперимент: {self.exp_config['name']}\")\n        if self.exp_config.get('pretrained_model_name'):\n            self.model = mobilenet_v2(weights=MobileNet_V2_Weights.IMAGENET1K_V1)\n            in_feat = self.model.classifier[1].in_features\n            self.model.classifier[1] = nn.Sequential(nn.Linear(in_feat, self.num_keypoints_coords), nn.Tanh())\n        else:\n            self.model = CustomCNNKeypoints(\n                num_keypoints_coords=self.num_keypoints_coords,\n                is_batch_norm=self.exp_config.get('batch_norm', True),\n                is_dropout=self.exp_config.get('dropout', True))\n        self.model = self.model.to(self.device)\n\n        opt_name = self.exp_config.get('optim', 'Adam'); lr = self.exp_config.get('lr', 1e-3)\n        if opt_name == 'SGD':\n            self.optimizer = optim.SGD(self.model.parameters(), lr=lr, momentum=0.9)\n        elif opt_name == 'RMSProp': \n            self.optimizer = optim.RMSprop(self.model.parameters(), lr=lr)\n        elif opt_name == 'Adam': \n            self.optimizer = optim.Adam(self.model.parameters(), lr=lr)\n            \n        self.scheduler = (lr_scheduler.ReduceLROnPlateau(self.optimizer, mode='min', factor=0.2, patience=3, verbose=True) \\\n                          if self.exp_config.get('sched', False) else None)\n\n    def _train_step(self, images, keypoints_true, valid_mask):\n        images, keypoints_true, valid_mask = images.to(self.device), keypoints_true.to(self.device), valid_mask.to(self.device)\n        self.optimizer.zero_grad()\n        keypoints_pred = self.model(images)\n        loss = self.criterion(keypoints_pred, keypoints_true, valid_mask) \n        loss.backward()\n        self.optimizer.step()\n        return loss.item(), keypoints_pred.detach(), keypoints_true.detach(), valid_mask.detach()\n\n    def _eval_step(self, images, keypoints_true, valid_mask):\n        images, keypoints_true, valid_mask = images.to(self.device), keypoints_true.to(self.device), valid_mask.to(self.device)\n        keypoints_pred = self.model(images)\n        loss = self.criterion(keypoints_pred, keypoints_true, valid_mask)\n        return loss.item(), keypoints_pred.detach(), keypoints_true.detach(), valid_mask.detach()\n\n    def _calculate_rmse_from_normalized(self, preds_norm, targets_norm, valid_mask, target_pixel_size=TARGET_SIZE):\n        preds_pixel = (preds_norm + 1.0) / 2.0 * target_pixel_size\n        targets_pixel = (targets_norm + 1.0) / 2.0 * target_pixel_size\n        \n        squared_errors = (preds_pixel - targets_pixel) ** 2\n        masked_squared_errors = squared_errors * valid_mask \n        \n        num_valid_coords = valid_mask.sum()\n        if num_valid_coords > 0:\n            mse_pixel = masked_squared_errors.sum() / num_valid_coords\n            return torch.sqrt(mse_pixel)\n        return torch.tensor(float('nan'), device=preds_norm.device)\n\n    def train_epoch(self):\n        self.model.train()\n        total_loss, batch_rmses, num_samples = 0, [], 0\n        for images, keypoints_true, valid_mask in tqdm(self.train_loader, desc=\"Training\", leave=False):\n            n = images.size(0)\n            loss, preds_norm, targets_norm, val_mask_batch = self._train_step(images, keypoints_true, valid_mask)\n            total_loss += loss * n; num_samples += n\n            rmse_val = self._calculate_rmse_from_normalized(preds_norm, targets_norm, val_mask_batch)\n            if not torch.isnan(rmse_val): batch_rmses.append(rmse_val.item())\n        avg_loss = total_loss / num_samples if num_samples else 0\n        avg_rmse = np.mean(batch_rmses) if batch_rmses else float('nan')\n        return avg_loss, avg_rmse\n\n    def eval_epoch(self, loader, desc=\"Evaluating\"):\n        self.model.eval()\n        total_loss, batch_rmses, num_samples = 0, [], 0\n        with torch.no_grad():\n            for images, keypoints_true, valid_mask in tqdm(loader, desc=desc, leave=False): \n                n = images.size(0)\n                loss, preds_norm, targets_norm, val_mask_batch = self._eval_step(images, keypoints_true, valid_mask) \n                total_loss += loss * n; num_samples += n\n                rmse_val = self._calculate_rmse_from_normalized(preds_norm, targets_norm, val_mask_batch)\n                if not torch.isnan(rmse_val): batch_rmses.append(rmse_val.item())\n        avg_loss = total_loss / num_samples if num_samples else 0\n        avg_rmse = np.mean(batch_rmses) if batch_rmses else float('nan')\n        return avg_loss, avg_rmse\n\n    def run(self): \n        history = {'train_loss': [], 'train_rmse': [], 'val_loss': [], 'val_rmse': []}\n        best_val_rmse = float('inf')\n\n        for epoch in range(1, self.num_epochs + 1):\n            train_loss, train_rmse = self.train_epoch()\n            val_loss, val_rmse = self.eval_epoch(self.valid_loader, desc=\"Validating\")\n            history['train_loss'].append(train_loss); history['train_rmse'].append(train_rmse)\n            history['val_loss'].append(val_loss); history['val_rmse'].append(val_rmse)\n            print(f\"Epoch {epoch}/{self.num_epochs}: tr_loss={train_loss:.4f}, tr_rmse={train_rmse:.2f}px | val_loss={val_loss:.4f}, val_rmse={val_rmse:.2f}px\")\n            if self.scheduler: self.scheduler.step(val_rmse)\n            if val_rmse < best_val_rmse:\n                best_val_rmse = val_rmse\n                print(f\"Новое лучшее val_rmse: {best_val_rmse:.2f}\")\n        return history, best_val_rmse","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:53:01.394782Z","iopub.execute_input":"2025-05-24T19:53:01.395301Z","iopub.status.idle":"2025-05-24T19:53:01.413932Z","shell.execute_reply.started":"2025-05-24T19:53:01.395277Z","shell.execute_reply":"2025-05-24T19:53:01.41313Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"exp_config_list = [\n    {\n        \"name\": \"MobileNetV2 + Adam\", \n        \"pretrained_model_name\": \"mobilenet_v2\", \n        \"optim\": \"Adam\",  \n        \"lr\": 0.0001,    \n        \"sched\": False,    \n        \"aug\": False,      \n\n    },\n    \n    {\n        \"name\": \"Custom CNN + SGD\",\n        \"optim\": \"SGD\",\n        \"lr\": 0.005,     \n        \"sched\": False,\n        \"aug\": False,\n        \"batch_norm\": True,\n        \"dropout\": True     \n    },\n\n    {\n        \"name\": \"Custom CNN + Adam\",\n        \"optim\": \"Adam\",\n        \"lr\": 0.0005,      \n        \"sched\": False,\n        \"aug\": False,\n        \"batch_norm\": True,\n        \"dropout\": True\n    },\n\n    {\n        \"name\": \"Custom CNN + Adam + Scheduler\",\n        \"optim\": \"Adam\",\n        \"lr\": 0.0005,\n        \"sched\": True,     \n        \"aug\": False,\n        \"batch_norm\": True,\n        \"dropout\": True\n    },\n\n    {\n        \"name\": \"Custom CNN + SGD + Augmentation\",\n        \"optim\": \"SGD\",    \n        \"lr\": 0.005,        \n        \"sched\": False,\n        \"aug\": True,      \n        \"batch_norm\": True,\n        \"dropout\": True\n    }\n]\n    \nall_histories = {}\nresults_table = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:53:01.457132Z","iopub.execute_input":"2025-05-24T19:53:01.457332Z","iopub.status.idle":"2025-05-24T19:53:01.462573Z","shell.execute_reply.started":"2025-05-24T19:53:01.457317Z","shell.execute_reply":"2025-05-24T19:53:01.461839Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for exp_conf in exp_config_list:\n    \n    trainer = TrainKeypointModels(\n        experiment_config=exp_conf, num_epochs=NUM_EPOCHS,\n        train_loader_no_aug=train_loader_without_aug, train_loader_with_aug=train_loader_with_aug,\n        valid_loader=valid_loader, CustomCNNKeypoints=CustomCNNKeypoints,\n        num_keypoints_coords=NUM_KEYPOINTS_COORDS, device=DEVICE, criterion=CRITERION\n    )\n    history, best_val_metric = trainer.run()\n    all_histories[exp_conf['name']] = history\n    results_table.append({\n        'Эксперимент': exp_conf['name'],\n        'Train RMSE (last epoch, px)': history['train_rmse'][-1] if history['train_rmse'] and not (isinstance(history['train_rmse'][-1], float) and np.isnan(history['train_rmse'][-1])) else 'N/A',\n        'Val RMSE (best, px)': best_val_metric if not (isinstance(best_val_metric, float) and np.isnan(best_val_metric)) else 'N/A'\n    })\n    print()\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:53:03.906079Z","iopub.execute_input":"2025-05-24T19:53:03.906596Z","iopub.status.idle":"2025-05-24T20:05:01.199992Z","shell.execute_reply.started":"2025-05-24T19:53:03.906571Z","shell.execute_reply":"2025-05-24T20:05:01.199166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results_df = pd.DataFrame(results_table)\nprint(\"\\n--- Сводная таблица результатов ---\")\nprint(results_df.to_string())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T20:05:01.201567Z","iopub.execute_input":"2025-05-24T20:05:01.201826Z","iopub.status.idle":"2025-05-24T20:05:01.210487Z","shell.execute_reply.started":"2025-05-24T20:05:01.201802Z","shell.execute_reply":"2025-05-24T20:05:01.209712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if results_table:\n    valid_results = [r for r in results_table if r['Val RMSE (best, px)'] != 'N/A']\n    if valid_results:\n        best_exp_result = min(valid_results, key=lambda x: x['Val RMSE (best, px)'])\n        best_exp_name = best_exp_result['Эксперимент']\n        best_history = all_histories.get(best_exp_name)\n        if best_history:\n            plt.figure(figsize=(12, 5)); plt.subplot(1, 2, 1)\n            plt.plot(best_history['train_loss'], label='Train Loss'); plt.plot(best_history['val_loss'], label='Validation Loss')\n            plt.xlabel('Epoch'); plt.ylabel('Loss (Masked MSE)'); plt.title(f'Loss for {best_exp_name}'); plt.legend(); plt.grid(True)\n            plt.subplot(1, 2, 2)\n            plt.plot(best_history['train_rmse'], label='Train RMSE (px)'); plt.plot(best_history['val_rmse'], label='Validation RMSE (px)')\n            plt.xlabel('Epoch'); plt.ylabel('RMSE (pixels)'); plt.title(f'RMSE for {best_exp_name}'); plt.legend(); plt.grid(True)\n            plt.tight_layout(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T20:05:07.816458Z","iopub.execute_input":"2025-05-24T20:05:07.817131Z","iopub.status.idle":"2025-05-24T20:05:08.189313Z","shell.execute_reply.started":"2025-05-24T20:05:07.817107Z","shell.execute_reply":"2025-05-24T20:05:08.188687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_predictions(model_to_viz, dataloader, device, num_samples=5, target_size=TARGET_SIZE, nan_placeholder_norm=None):\n    model_to_viz.eval(); samples_shown = 0\n    data_iter = iter(dataloader)\n    plt.figure(figsize=(15, num_samples * 4)); plot_idx = 1\n\n    if nan_placeholder_norm is None: \n         nan_placeholder_norm = (NAN_PLACEHOLDER / float(target_size)) * 2.0 - 1.0\n\n\n    with torch.no_grad():\n        for _ in range(min(len(dataloader), num_samples * BATCH_SIZE)): \n            if samples_shown >= num_samples: break\n            try: images_batch, keypoints_true_batch_norm, valid_mask_batch = next(data_iter)\n            except StopIteration: break\n\n            images_batch_dev = images_batch.to(device)\n            keypoints_pred_batch_norm = model_to_viz(images_batch_dev).cpu()\n\n            for i in range(images_batch.size(0)):\n                if samples_shown >= num_samples: break\n                img_tensor = images_batch[i].cpu()\n                true_kps_norm = keypoints_true_batch_norm[i]\n                pred_kps_norm = keypoints_pred_batch_norm[i]\n                current_valid_mask = valid_mask_batch[i].cpu().numpy() \n\n                inv_normalize = transforms.Normalize(mean=[-0.485/0.229, -0.456/0.224, -0.406/0.225], std=[1/0.229, 1/0.224, 1/0.225])\n                img_display = inv_normalize(img_tensor).permute(1, 2, 0).numpy(); img_display = np.clip(img_display, 0, 1)\n                \n                true_kps_pixel = ((true_kps_norm.numpy() + 1.0) / 2.0) * target_size\n                pred_kps_pixel = ((pred_kps_norm.numpy() + 1.0) / 2.0) * target_size\n                \n                plt.subplot(num_samples, 2, plot_idx); plt.imshow(img_display); plt.title(\"True Keypoints\"); plt.axis('off')\n                for k_idx in range(0, len(true_kps_pixel), 2):\n                    if current_valid_mask[k_idx] > 0.5 and current_valid_mask[k_idx+1] > 0.5: \n                        plt.scatter(true_kps_pixel[k_idx], true_kps_pixel[k_idx+1], s=20, marker='o', c='cyan', edgecolors='black', linewidths=0.5)\n                plot_idx += 1\n\n                plt.subplot(num_samples, 2, plot_idx); plt.imshow(img_display); plt.title(\"Predicted Keypoints\"); plt.axis('off')\n                for k_idx in range(0, len(pred_kps_pixel), 2):\n                    if current_valid_mask[k_idx] > 0.5 and current_valid_mask[k_idx+1] > 0.5: \n                        plt.scatter(pred_kps_pixel[k_idx], pred_kps_pixel[k_idx+1], s=20, marker='x', c='red')\n                plot_idx += 1; samples_shown += 1\n    plt.tight_layout(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T20:05:13.463251Z","iopub.execute_input":"2025-05-24T20:05:13.463904Z","iopub.status.idle":"2025-05-24T20:05:13.474574Z","shell.execute_reply.started":"2025-05-24T20:05:13.463877Z","shell.execute_reply":"2025-05-24T20:05:13.473769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 'trainer' in locals() and hasattr(trainer, 'model'):\n    print(f\"\\n--- Визуализация предсказаний для последней обученной модели: {trainer.exp_config['name']} ---\")\n    trainer.model.to(DEVICE) \n    visualize_predictions(trainer.model, valid_loader, DEVICE, num_samples=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T20:06:04.972361Z","iopub.execute_input":"2025-05-24T20:06:04.973126Z","iopub.status.idle":"2025-05-24T20:06:07.307759Z","shell.execute_reply.started":"2025-05-24T20:06:04.973103Z","shell.execute_reply":"2025-05-24T20:06:07.306885Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Задание 2 CNN локализация лицевых точек","metadata":{}},{"cell_type":"markdown","source":"https://www.kaggle.com/datasets/bulentsiyah/semantic-drone-dataset","metadata":{}},{"cell_type":"code","source":"!pip install segmentation_models_pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:01:35.585966Z","iopub.execute_input":"2025-05-24T17:01:35.586595Z","iopub.status.idle":"2025-05-24T17:02:52.356215Z","shell.execute_reply.started":"2025-05-24T17:01:35.586571Z","shell.execute_reply":"2025-05-24T17:02:52.355235Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nfrom tqdm.notebook import tqdm \nimport cv2\nimport segmentation_models_pytorch as smp","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:32.911196Z","iopub.execute_input":"2025-05-24T17:20:32.911499Z","iopub.status.idle":"2025-05-24T17:20:32.915845Z","shell.execute_reply.started":"2025-05-24T17:20:32.911476Z","shell.execute_reply":"2025-05-24T17:20:32.914943Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMAGE_SIZE = 128\nBATCH_SIZE = 8\nEPOCHS = 5\n\nIMAGE_DIR = '/kaggle/input/semantic-drone-dataset/dataset/semantic_drone_dataset/original_images/'\nMASK_DIR = '/kaggle/input/semantic-drone-dataset/dataset/semantic_drone_dataset/label_images_semantic/'\nCLASS_DICT_PATH = '/kaggle/input/semantic-drone-dataset/class_dict_seg.csv'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:33.202497Z","iopub.execute_input":"2025-05-24T17:20:33.202784Z","iopub.status.idle":"2025-05-24T17:20:33.207009Z","shell.execute_reply.started":"2025-05-24T17:20:33.202764Z","shell.execute_reply":"2025-05-24T17:20:33.206433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def parse_class_dict(csv_path):\n    df = pd.read_csv(csv_path)\n    df.columns = df.columns.str.strip() \n    \n    class_names = df['name'].tolist()\n    class_rgb_values = df[['r', 'g', 'b']].values.tolist()\n\n    color_to_id_map = {tuple(rgb): i for i, rgb in enumerate(class_rgb_values)}\n    id_to_color_map = {i: tuple(rgb) for i, rgb in enumerate(class_rgb_values)}\n    num_classes = len(class_names)\n    \n    return class_names, class_rgb_values, color_to_id_map, id_to_color_map, num_classes\n\nCLASS_NAMES, CLASS_RGB_VALUES, COLOR_TO_ID, ID_TO_COLOR, NUM_CLASSES = parse_class_dict(CLASS_DICT_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:34.655854Z","iopub.execute_input":"2025-05-24T17:20:34.656517Z","iopub.status.idle":"2025-05-24T17:20:34.668188Z","shell.execute_reply.started":"2025-05-24T17:20:34.656494Z","shell.execute_reply":"2025-05-24T17:20:34.667517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rgb_to_mask(rgb_image_np, color_to_id_map):\n    h, w, _ = rgb_image_np.shape\n    mask = np.zeros((h, w), dtype=np.int64) \n    for rgb_tuple, class_id in color_to_id_map.items():\n        condition = (rgb_image_np[..., 0] == rgb_tuple[0]) & \\\n                    (rgb_image_np[..., 1] == rgb_tuple[1]) & \\\n                    (rgb_image_np[..., 2] == rgb_tuple[2])\n        mask[condition] = class_id\n    return mask\n\ndef mask_to_rgb(mask_np, id_to_color_map):\n    h, w = mask_np.shape\n    rgb_mask = np.zeros((h, w, 3), dtype=np.uint8)\n    for class_id, rgb_tuple in id_to_color_map.items():\n        rgb_mask[mask_np == class_id] = rgb_tuple\n    return rgb_mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:34.835027Z","iopub.execute_input":"2025-05-24T17:20:34.835222Z","iopub.status.idle":"2025-05-24T17:20:34.840645Z","shell.execute_reply.started":"2025-05-24T17:20:34.835208Z","shell.execute_reply":"2025-05-24T17:20:34.840036Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CustomAerialDataset(Dataset):\n    def __init__(self, image_paths, mask_paths, color_to_id_map, augmentation=None, preprocessing=None):\n        self.image_paths = image_paths\n        self.mask_paths = mask_paths\n        self.color_to_id_map = color_to_id_map\n        self.augmentation = augmentation\n        self.preprocessing = preprocessing\n        \n        assert len(self.image_paths) == len(self.mask_paths)\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image_path = self.image_paths[idx]\n        mask_path = self.mask_paths[idx]\n\n        image = cv2.imread(image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        mask_rgb = cv2.imread(mask_path)\n        mask_rgb = cv2.cvtColor(mask_rgb, cv2.COLOR_BGR2RGB)\n        \n        mask = rgb_to_mask(mask_rgb, self.color_to_id_map)\n\n        if self.augmentation:\n            augmented = self.augmentation(image=image, mask=mask)\n            image = augmented['image']\n            mask = augmented['mask']\n        \n        if self.preprocessing:\n            preprocessed = self.preprocessing(image=image)\n            image = preprocessed['image']\n\n        image = image.float()\n        mask = torch.from_numpy(mask).long()\n\n        return image, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:36.481035Z","iopub.execute_input":"2025-05-24T17:20:36.481827Z","iopub.status.idle":"2025-05-24T17:20:36.489206Z","shell.execute_reply.started":"2025-05-24T17:20:36.481792Z","shell.execute_reply":"2025-05-24T17:20:36.488401Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mask_filenames_only = sorted([f for f in os.listdir(MASK_DIR) if f.endswith('.png')])\n\nall_image_full_paths = []\nall_mask_full_paths = []\n\nfor mask_fn_only in mask_filenames_only:\n    image_fn_only = mask_fn_only.replace('.png', '.jpg')\n    \n    current_image_path = os.path.join(IMAGE_DIR, image_fn_only)\n    current_mask_path = os.path.join(MASK_DIR, mask_fn_only)\n    \n    if os.path.exists(current_image_path):\n        all_image_full_paths.append(current_image_path)\n        all_mask_full_paths.append(current_mask_path)\n\nprint(f\"Найдено {len(all_image_full_paths)} согласованных пар изображение/маска.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:36.680121Z","iopub.execute_input":"2025-05-24T17:20:36.680307Z","iopub.status.idle":"2025-05-24T17:20:36.993036Z","shell.execute_reply.started":"2025-05-24T17:20:36.680293Z","shell.execute_reply":"2025-05-24T17:20:36.992268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_image_paths, val_image_paths, train_mask_paths, val_mask_paths = train_test_split(\n    all_image_full_paths, all_mask_full_paths, test_size=0.2, random_state=RANDOM_STATE\n)\n\nprint(f\"Количество обучающих изображений: {len(train_image_paths)}\")\nprint(f\"Количество валидационных изображений: {len(val_image_paths)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:36.994185Z","iopub.execute_input":"2025-05-24T17:20:36.994527Z","iopub.status.idle":"2025-05-24T17:20:37.000564Z","shell.execute_reply.started":"2025-05-24T17:20:36.994505Z","shell.execute_reply":"2025-05-24T17:20:36.999987Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_train_augs(img_size):\n    return A.Compose([\n        A.Resize(img_size, img_size, interpolation=cv2.INTER_NEAREST),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n    ])\n\ndef get_val_augs(img_size):\n    return A.Compose([\n        A.Resize(img_size, img_size, interpolation=cv2.INTER_NEAREST),\n    ])\n\n\ndef get_smp_preprocessing(preprocessing_fn):\n    _transform = [\n        A.Lambda(image=preprocessing_fn),\n        A.ToTensorV2(),\n    ]\n    return A.Compose(_transform)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:38.539524Z","iopub.execute_input":"2025-05-24T17:20:38.540108Z","iopub.status.idle":"2025-05-24T17:20:38.545037Z","shell.execute_reply.started":"2025-05-24T17:20:38.540086Z","shell.execute_reply":"2025-05-24T17:20:38.544212Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, loss_fn, device):\n    model.train()\n    epoch_loss = 0.0\n    \n    all_ious = []\n    progress_bar = tqdm(loader, desc=\"Training\", leave=False)\n    for images, masks in progress_bar:\n        images = images.to(device)\n        masks = masks.to(device).long()\n\n        optimizer.zero_grad()\n        outputs = model(images)\n        \n        loss = loss_fn(outputs, masks)\n        loss.backward()\n        optimizer.step()\n\n        epoch_loss += loss.item()\n        pred_masks_classes = torch.argmax(outputs, dim=1) \n        \n        tp, fp, fn, tn = smp.metrics.get_stats(\n            output=pred_masks_classes, \n            target=masks, \n            mode='multiclass', \n            num_classes=NUM_CLASSES\n        )\n        \n        iou_val = smp.metrics.iou_score(tp, fp, fn, tn, reduction='micro-imagewise') \n        \n        all_ious.append(iou_val.item())\n\n        progress_bar.set_postfix(loss=loss.item(), iou=iou_val.item())\n        \n    avg_loss = epoch_loss / len(loader)\n    avg_iou = np.mean(all_ious) if all_ious else 0.0\n    return avg_loss, avg_iou","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:40.484372Z","iopub.execute_input":"2025-05-24T17:20:40.484671Z","iopub.status.idle":"2025-05-24T17:20:40.491032Z","shell.execute_reply.started":"2025-05-24T17:20:40.484651Z","shell.execute_reply":"2025-05-24T17:20:40.490199Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate_one_epoch(model, loader, loss_fn, device):\n    model.eval()\n    epoch_loss = 0.0\n    all_ious = []\n\n    progress_bar = tqdm(loader, desc=\"Validation\", leave=False)\n    with torch.no_grad():\n        for images, masks in progress_bar:\n            images = images.to(device)\n            masks = masks.to(device).long() \n\n            outputs = model(images)\n            loss = loss_fn(outputs, masks)\n            epoch_loss += loss.item()\n\n            pred_masks_classes = torch.argmax(outputs, dim=1)\n            \n            tp, fp, fn, tn = smp.metrics.get_stats(\n                output=pred_masks_classes, \n                target=masks, \n                mode='multiclass', \n                num_classes=NUM_CLASSES\n            )\n            \n            iou_val = smp.metrics.iou_score(tp, fp, fn, tn, reduction='micro-imagewise')\n            \n            all_ious.append(iou_val.item())\n            progress_bar.set_postfix(loss=loss.item(), iou=iou_val.item())\n\n    avg_loss = epoch_loss / len(loader)\n    avg_iou = np.mean(all_ious) if all_ious else 0.0\n    return avg_loss, avg_iou","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:40.715717Z","iopub.execute_input":"2025-05-24T17:20:40.716281Z","iopub.status.idle":"2025-05-24T17:20:40.72228Z","shell.execute_reply.started":"2025-05-24T17:20:40.716257Z","shell.execute_reply":"2025-05-24T17:20:40.721479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models_to_train_config = [\n    {\n        \"name\": \"Unet_ResNet34\",\n        \"arch_fn\": smp.Unet,\n        \"encoder\": \"resnet34\",\n        \"encoder_weights\": \"imagenet\",\n    },\n    {\n        \"name\": \"DeepLabV3Plus_MobileNetV2\",\n        \"arch_fn\": smp.DeepLabV3Plus,\n        \"encoder\": \"mobilenet_v2\",\n        \"encoder_weights\": \"imagenet\",\n    },\n\n    {\n        \"name\": \"FPN_EfficientNetB0\",\n        \"arch_fn\": smp.FPN,\n        \"encoder\": \"efficientnet-b0\",\n        \"encoder_weights\": \"imagenet\",\n    }\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:42.203406Z","iopub.execute_input":"2025-05-24T17:20:42.204112Z","iopub.status.idle":"2025-05-24T17:20:42.20845Z","shell.execute_reply.started":"2025-05-24T17:20:42.204087Z","shell.execute_reply":"2025-05-24T17:20:42.207648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results_comparison = []\n\nfor model_config in models_to_train_config:\n    print(f\"Обучение модели: {model_config['name']}\")\n    \n    smp_preprocessing_fn = smp.encoders.get_preprocessing_fn(\n        model_config['encoder'], \n        pretrained=model_config['encoder_weights']\n    )\n\n    train_dataset = CustomAerialDataset(\n        image_paths=train_image_paths,\n        mask_paths=train_mask_paths,\n        color_to_id_map=COLOR_TO_ID,\n        augmentation=get_train_augs(IMAGE_SIZE),\n        preprocessing=get_smp_preprocessing(smp_preprocessing_fn)\n    )\n\n    val_dataset = CustomAerialDataset(\n        image_paths=val_image_paths,\n        mask_paths=val_mask_paths,\n        color_to_id_map=COLOR_TO_ID,\n        augmentation=get_val_augs(IMAGE_SIZE),\n        preprocessing=get_smp_preprocessing(smp_preprocessing_fn)\n    )\n    \n    train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2, pin_memory=True)\n    val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2, pin_memory=True)\n\n    model = model_config['arch_fn'](\n        encoder_name=model_config['encoder'],\n        encoder_weights=model_config['encoder_weights'],\n        classes=NUM_CLASSES,\n        activation=None,\n    )\n    model.to(DEVICE)\n\n    loss_fn = smp.losses.DiceLoss(mode='multiclass', from_logits=True)\n    optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n\n    best_val_iou = 0.0\n    model_save_path = f\"{model_config['name']}_best_model.pth\"\n\n    for epoch in range(EPOCHS):\n        print(f\"Эпоха {epoch+1}/{EPOCHS}\")\n        train_loss, train_iou = train_one_epoch(model, train_loader, optimizer, loss_fn, DEVICE)\n        val_loss, val_iou = validate_one_epoch(model, val_loader, loss_fn, DEVICE)\n        \n        print(f\"Обучение: Loss: {train_loss:.4f}, IoU: {train_iou:.4f}\")\n        print(f\"Валидация: Loss: {val_loss:.4f}, IoU: {val_iou:.4f}\")\n\n        if val_iou > best_val_iou:\n            best_val_iou = val_iou\n            torch.save(model.state_dict(), model_save_path)\n            print(f\"Модель сохранена: {model_save_path} (Val IoU: {best_val_iou:.4f})\")\n    \n    results_comparison.append({\n        \"model_name\": model_config['name'],\n        \"best_val_iou\": best_val_iou,\n        \"saved_path\": model_save_path if best_val_iou > 0 else None\n    })","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T17:20:42.454647Z","iopub.execute_input":"2025-05-24T17:20:42.455009Z","iopub.status.idle":"2025-05-24T19:15:26.673695Z","shell.execute_reply.started":"2025-05-24T17:20:42.454992Z","shell.execute_reply":"2025-05-24T19:15:26.67291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"comparison_df = pd.DataFrame(results_comparison)\nprint(comparison_df.to_string())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:15:26.908951Z","iopub.execute_input":"2025-05-24T19:15:26.909169Z","iopub.status.idle":"2025-05-24T19:15:26.915102Z","shell.execute_reply.started":"2025-05-24T19:15:26.909145Z","shell.execute_reply":"2025-05-24T19:15:26.914453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_predictions(\n    val_img_paths, \n    val_msk_paths, \n    trained_models_info, \n    id_to_color_map_viz, \n    num_samples=3, \n    img_size_viz=IMAGE_SIZE\n    ):\n        \n    indices_to_show = random.sample(range(len(val_img_paths)), min(num_samples, len(val_img_paths)))\n\n    for i in indices_to_show:\n        original_image_path = val_img_paths[i]\n        gt_mask_path = val_msk_paths[i]\n\n        img_for_display_raw = cv2.imread(original_image_path)\n        img_for_display_raw = cv2.cvtColor(img_for_display_raw, cv2.COLOR_BGR2RGB)\n        \n        gt_mask_rgb_raw = cv2.imread(gt_mask_path)\n        gt_mask_rgb_raw = cv2.cvtColor(gt_mask_rgb_raw, cv2.COLOR_BGR2RGB)\n\n        resize_aug = A.Resize(img_size_viz, img_size_viz, interpolation=cv2.INTER_NEAREST)\n        \n        resized_img_for_display = resize_aug(image=img_for_display_raw)['image']\n        resized_gt_mask_for_display = resize_aug(image=gt_mask_rgb_raw)['image'] \n\n        num_cols = 2 + len(trained_models_info) \n        plt.figure(figsize=(5 * num_cols, 5))\n\n        plt.subplot(1, num_cols, 1)\n        plt.imshow(resized_img_for_display)\n        plt.title(\"Original Image\")\n        plt.axis(\"off\")\n\n        plt.subplot(1, num_cols, 2)\n        plt.imshow(resized_gt_mask_for_display)\n        plt.title(\"Ground Truth Mask\")\n        plt.axis(\"off\")\n\n        plot_idx = 3\n        for model_res in trained_models_info:\n            plt.subplot(1, num_cols, plot_idx)\n            if model_res[\"saved_path\"] is None or not os.path.exists(model_res[\"saved_path\"]):\n                print(f\"Модель {model_res['model_name']} не обучена или файл весов не найден.\")\n                plt.title(f\"{model_res['model_name']}\\n(Not Available)\")\n                plt.axis(\"off\")\n                plot_idx += 1\n                continue\n\n            current_model_config = next(m for m in models_to_train_config if m['name'] == model_res['model_name'])\n \n            model_viz = current_model_config['arch_fn'](\n                encoder_name=current_model_config['encoder'],\n                classes=NUM_CLASSES,\n                activation=None\n            )\n            model_viz.load_state_dict(torch.load(model_res[\"saved_path\"], map_location=DEVICE))\n            model_viz.to(DEVICE)\n            model_viz.eval()\n\n            smp_preprocessing_fn_viz = smp.encoders.get_preprocessing_fn(\n                current_model_config['encoder'], \n                pretrained=current_model_config['encoder_weights']\n            )\n            smp_processor_viz = get_smp_preprocessing(smp_preprocessing_fn_viz)\n            \n            input_for_model = resize_aug(image=img_for_display_raw)['image'] \n            input_tensor_for_pred = smp_processor_viz(image=input_for_model)['image']\n            input_tensor_for_pred = input_tensor_for_pred.unsqueeze(0).to(DEVICE).float()\n\n            with torch.no_grad():\n                pred_logits = model_viz(input_tensor_for_pred)\n                pred_mask_indices = torch.argmax(pred_logits, dim=1).squeeze(0).cpu().numpy()\n\n            pred_mask_rgb = mask_to_rgb(pred_mask_indices, id_to_color_map_viz)\n\n            plt.imshow(pred_mask_rgb)\n            plt.title(f\"{model_res['model_name']}\\nPred (IoU: {model_res['best_val_iou']:.3f})\")\n            plt.axis(\"off\")\n            plot_idx += 1\n        \n        plt.tight_layout()\n        plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:18:55.569348Z","iopub.execute_input":"2025-05-24T19:18:55.570187Z","iopub.status.idle":"2025-05-24T19:18:55.580538Z","shell.execute_reply.started":"2025-05-24T19:18:55.57015Z","shell.execute_reply":"2025-05-24T19:18:55.579777Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if val_image_paths: # Убедимся, что есть что визуализировать\n    visualize_predictions(\n        val_image_paths, \n        val_mask_paths, \n        results_comparison, # DataFrame с результатами моделей\n        ID_TO_COLOR, \n        num_samples=3, # Количество примеров для отображения\n        img_size_viz=IMAGE_SIZE # Размер для визуализации\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T19:18:56.707051Z","iopub.execute_input":"2025-05-24T19:18:56.707665Z","iopub.status.idle":"2025-05-24T19:19:02.311407Z","shell.execute_reply.started":"2025-05-24T19:18:56.70764Z","shell.execute_reply":"2025-05-24T19:19:02.310751Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}