{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":132097,"databundleVersionId":15841209}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Импорт библиотек","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\nimport os\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport torchvision\nfrom torchvision.datasets import ImageFolder\n\nfrom torch.utils.data.dataloader import DataLoader\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import Subset\n\nimport matplotlib.pyplot as plt\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport random\nfrom PIL import Image\n\n%matplotlib inline\ntorch.manual_seed(0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T11:58:10.270738Z","iopub.execute_input":"2026-02-28T11:58:10.2715Z","iopub.status.idle":"2026-02-28T11:58:24.146259Z","shell.execute_reply.started":"2026-02-28T11:58:10.27147Z","shell.execute_reply":"2026-02-28T11:58:24.145209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Проверка доступа к GPU","metadata":{}},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device=torch.device(\"cuda:0\")\n    print(\"Training on GPU...\")\nelse:\n    device = torch.device(\"cpu\")\n    print(\"Training on CPU...\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T11:58:24.147464Z","iopub.execute_input":"2026-02-28T11:58:24.147823Z","iopub.status.idle":"2026-02-28T11:58:24.198232Z","shell.execute_reply.started":"2026-02-28T11:58:24.147801Z","shell.execute_reply":"2026-02-28T11:58:24.1975Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Преобразования (Transforms)","metadata":{}},{"cell_type":"code","source":"# Преобразования для обучающей выборки с аугментацией и нормализацией\ntrain_transform = torchvision.transforms.Compose([\n    torchvision.transforms.Resize(size=(224, 224)),\n    torchvision.transforms.RandomHorizontalFlip(),\n    torchvision.transforms.RandomRotation(10),\n    torchvision.transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),\n    torchvision.transforms.ToTensor(),\n    torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                     std=[0.229, 0.224, 0.225])\n])\n\n# Преобразования для валидации и теста (только ресайз, тензор, нормализация)\nval_transform = torchvision.transforms.Compose([\n    torchvision.transforms.Resize(size=(224, 224)),\n    torchvision.transforms.ToTensor(),\n    torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                     std=[0.229, 0.224, 0.225])\n])\n\ntest_transform = val_transform  # одинаково для теста","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T11:58:24.199241Z","iopub.execute_input":"2026-02-28T11:58:24.199494Z","iopub.status.idle":"2026-02-28T11:58:24.215134Z","shell.execute_reply.started":"2026-02-28T11:58:24.199475Z","shell.execute_reply":"2026-02-28T11:58:24.214512Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Создание объектов Datasets с разбиением выборки на обучающую и валидационную","metadata":{}},{"cell_type":"code","source":"train_val_path=\"/kaggle/input/optical-coherence-tomography-classification/Dataset\"\n\ntrain_dataset = ImageFolder(train_val_path, transform=train_transform)\nval_dataset = ImageFolder(train_val_path, transform=val_transform)\n\nclass_names = train_dataset.classes\nprint(class_names) # list out all the classes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T11:58:24.216703Z","iopub.execute_input":"2026-02-28T11:58:24.216971Z","iopub.status.idle":"2026-02-28T12:01:50.283959Z","shell.execute_reply.started":"2026-02-28T11:58:24.216953Z","shell.execute_reply":"2026-02-28T12:01:50.283119Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Splitting the data into train and validation set\ndef split_train_val(tot_img, val_percentage=0.2, rnd=23):\n    # Here indices are randomly permuted \n    number_of_val = int(tot_img*val_percentage)\n    \n    np.random.seed(rnd)\n    indexs = np.random.permutation(tot_img)\n    return indexs[0:number_of_val], indexs[number_of_val:]\n\nrandomness = 1\nval_per = 0.2\n\nall_len = len(train_dataset)\n\nval_indices, train_indices = split_train_val(all_len, val_per, randomness)\n\nprint(val_indices, \"validation data:\", val_indices.shape)\nprint(train_indices, \"train data:\", train_indices.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:01:50.285101Z","iopub.execute_input":"2026-02-28T12:01:50.285519Z","iopub.status.idle":"2026-02-28T12:01:50.294159Z","shell.execute_reply.started":"2026-02-28T12:01:50.285488Z","shell.execute_reply":"2026-02-28T12:01:50.2936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = Subset(train_dataset, train_indices)\nval_dataset = Subset(val_dataset, val_indices)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:01:50.295253Z","iopub.execute_input":"2026-02-28T12:01:50.29578Z","iopub.status.idle":"2026-02-28T12:01:50.308693Z","shell.execute_reply.started":"2026-02-28T12:01:50.29576Z","shell.execute_reply":"2026-02-28T12:01:50.308207Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img0, label0 = train_dataset[9627]\nprint(img0.shape, label0)\n\nimg1,label1 = val_dataset[20]\nprint(img1.shape, label1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:01:50.309441Z","iopub.execute_input":"2026-02-28T12:01:50.30966Z","iopub.status.idle":"2026-02-28T12:01:50.541651Z","shell.execute_reply.started":"2026-02-28T12:01:50.309634Z","shell.execute_reply":"2026-02-28T12:01:50.540835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show(img, label):\n    print(\"label-->\",class_names[label])\n    img = img.numpy().transpose((1, 2, 0)) # Channel first then height and width\n    # Если вы применили нормализацию, например, для ResNet18,\n    # то для визуализации и корректного отображения нужно сделать обратные преобразования\n    # mean = np.array([0.485, 0.456, 0.406])\n    # std = np.array([0.229, 0.224, 0.225])\n    # img = img * std + mean\n    # img = np.clip(img, 0., 1.)\n    plt.imshow(img)\n\nshow(*train_dataset[6])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:01:50.542621Z","iopub.execute_input":"2026-02-28T12:01:50.542925Z","iopub.status.idle":"2026-02-28T12:01:50.720546Z","shell.execute_reply.started":"2026-02-28T12:01:50.542902Z","shell.execute_reply":"2026-02-28T12:01:50.719947Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Dataloaders","metadata":{}},{"cell_type":"code","source":"batch_size = 256\n\ntrain_dataloader = DataLoader(train_dataset, batch_size, shuffle=True)\nval_dataloader = DataLoader(val_dataset, batch_size, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:01:50.721334Z","iopub.execute_input":"2026-02-28T12:01:50.721525Z","iopub.status.idle":"2026-02-28T12:01:50.725661Z","shell.execute_reply.started":"2026-02-28T12:01:50.721508Z","shell.execute_reply":"2026-02-28T12:01:50.724959Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Создание модели","metadata":{}},{"cell_type":"code","source":"import torchvision.models as models\n\n# Загружаем предобученный ResNet18\nmodel = models.resnet18(pretrained=True)\n\n# Заменяем последний fully-connected слой на 4 класса\nnum_ftrs = model.fc.in_features\nmodel.fc = nn.Linear(num_ftrs, 4)\n\nprint(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:01:50.727905Z","iopub.execute_input":"2026-02-28T12:01:50.728169Z","iopub.status.idle":"2026-02-28T12:01:51.28083Z","shell.execute_reply.started":"2026-02-28T12:01:50.728142Z","shell.execute_reply":"2026-02-28T12:01:51.280214Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"class_names","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:04:11.606822Z","iopub.execute_input":"2026-02-28T12:04:11.607396Z","iopub.status.idle":"2026-02-28T12:04:11.611655Z","shell.execute_reply.started":"2026-02-28T12:04:11.60737Z","shell.execute_reply":"2026-02-28T12:04:11.611026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = model.to(device)\n\ncriterion = nn.CrossEntropyLoss()\n\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n\n# Планировщик: уменьшаем lr в 10 раз, если val_loss перестал уменьшаться\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:05:30.071538Z","iopub.execute_input":"2026-02-28T12:05:30.07183Z","iopub.status.idle":"2026-02-28T12:05:30.078216Z","shell.execute_reply.started":"2026-02-28T12:05:30.071809Z","shell.execute_reply":"2026-02-28T12:05:30.077459Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(epochs):\n    print('Starting training..')\n    history = {'train_loss': [], 'val_loss': [], 'val_acc': []}\n    \n    for epoch in range(epochs):\n        print('='*20)\n        print(f'Starting epoch {epoch + 1}/{epochs}')\n        print('='*20)\n\n        # Обучение\n        model.train()\n        train_loss = 0.0\n        for images, labels in train_dataloader:\n            images, labels = images.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n\n            train_loss += loss.item()\n\n        train_loss /= len(train_dataloader)\n        history['train_loss'].append(train_loss)\n\n        # Валидация\n        model.eval()\n        val_loss = 0.0\n        correct = 0\n        total = 0\n\n        with torch.no_grad():\n            for images, labels in val_dataloader:\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                val_loss += loss.item()\n\n                _, predicted = torch.max(outputs, 1)\n                total += labels.size(0)\n                correct += (predicted == labels).sum().item()\n\n        val_loss /= len(val_dataloader)\n        accuracy = correct / total\n\n        history['val_loss'].append(val_loss)\n        history['val_acc'].append(accuracy)\n\n        print(f'Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | Val Acc: {accuracy:.4f}')\n\n        # Планировщик шагает по val_loss\n        scheduler.step(val_loss)\n\n        # Ранняя остановка, если accuracy > 0.95 (не обязательно)\n        if accuracy > 0.5:\n            print(\"Target accuracy reached, stopping.\")\n            break\n\n    print('Training complete.')\n    return history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:05:37.125195Z","iopub.execute_input":"2026-02-28T12:05:37.125693Z","iopub.status.idle":"2026-02-28T12:05:37.133198Z","shell.execute_reply.started":"2026-02-28T12:05:37.125668Z","shell.execute_reply":"2026-02-28T12:05:37.132638Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Запуск обучения","metadata":{}},{"cell_type":"code","source":"%%time\n\ntrain(epochs=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:05:56.210726Z","iopub.execute_input":"2026-02-28T12:05:56.211043Z","iopub.status.idle":"2026-02-28T12:30:37.742654Z","shell.execute_reply.started":"2026-02-28T12:05:56.211013Z","shell.execute_reply":"2026-02-28T12:30:37.741792Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nhistory = train(epochs=10)  # запускаем обучение на 10 эпох\n\n# График потерь\nplt.figure(figsize=(12, 4))\nplt.subplot(1, 2, 1)\nplt.plot(history['train_loss'], label='Train Loss')\nplt.plot(history['val_loss'], label='Val Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.title('Loss over epochs')\n\n# График точности\nplt.subplot(1, 2, 2)\nplt.plot(history['val_acc'], label='Val Accuracy', color='green')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.legend()\nplt.title('Validation accuracy over epochs')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:33:13.176461Z","iopub.execute_input":"2026-02-28T12:33:13.176803Z","iopub.status.idle":"2026-02-28T12:49:13.678771Z","shell.execute_reply.started":"2026-02-28T12:33:13.176779Z","shell.execute_reply":"2026-02-28T12:49:13.678076Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay\nimport numpy as np\n\nmodel.eval()\nall_preds = []\nall_labels = []\n\nwith torch.no_grad():\n    for images, labels in val_dataloader:\n        images, labels = images.to(device), labels.to(device)\n        outputs = model(images)\n        _, preds = torch.max(outputs, 1)\n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\ncm = confusion_matrix(all_labels, all_preds)\ndisp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=class_names)\ndisp.plot(cmap=plt.cm.Blues)\nplt.title('Confusion Matrix on Validation Set')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:50:47.352189Z","iopub.execute_input":"2026-02-28T12:50:47.352912Z","iopub.status.idle":"2026-02-28T12:52:55.87095Z","shell.execute_reply.started":"2026-02-28T12:50:47.352885Z","shell.execute_reply":"2026-02-28T12:52:55.870156Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Создание класса TestDataset\nНаследуемся от базового класса Dataset, модифицируем метод get_item, который теперь возвращает только изображение, т.к. метка класса неизвестна и ее нужно предсказать в рамках соревнования","metadata":{}},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, root_dir, transform=None): \n        self.root_dir = root_dir\n        self.transform = transform\n        self.filenames = sorted(os.listdir(root_dir))\n\n    def __len__(self):\n        return len(self.filenames)\n\n    def __getitem__(self, idx):\n        img_name = os.path.join(self.root_dir,\n                                self.filenames[idx])\n        \n        image = Image.open(img_name)\n        image = image.convert('RGB')\n        \n        if self.transform:\n            sample = self.transform(image)\n\n        return sample","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:55:43.070365Z","iopub.execute_input":"2026-02-28T12:55:43.07063Z","iopub.status.idle":"2026-02-28T12:55:43.075663Z","shell.execute_reply.started":"2026-02-28T12:55:43.07061Z","shell.execute_reply":"2026-02-28T12:55:43.074909Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"root_dir = '/kaggle/input/optical-coherence-tomography-classification/Test'\ntest_dataset = TestDataset(root_dir, test_transform)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:55:47.34591Z","iopub.execute_input":"2026-02-28T12:55:47.346763Z","iopub.status.idle":"2026-02-28T12:55:47.350965Z","shell.execute_reply.started":"2026-02-28T12:55:47.346725Z","shell.execute_reply":"2026-02-28T12:55:47.350442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img0 = test_dataset[112]\nprint(img0.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:55:49.781012Z","iopub.execute_input":"2026-02-28T12:55:49.781739Z","iopub.status.idle":"2026-02-28T12:55:49.795025Z","shell.execute_reply.started":"2026-02-28T12:55:49.781711Z","shell.execute_reply":"2026-02-28T12:55:49.794298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataloader = DataLoader(test_dataset, batch_size, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:55:52.343786Z","iopub.execute_input":"2026-02-28T12:55:52.344386Z","iopub.status.idle":"2026-02-28T12:55:52.347828Z","shell.execute_reply.started":"2026-02-28T12:55:52.344358Z","shell.execute_reply":"2026-02-28T12:55:52.347245Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Predict на тестовой выборке","metadata":{}},{"cell_type":"code","source":"# Generate predictions\npredictions = []\n\nmodel.eval()  # Set model to evaluation mode\n\nfor images in test_dataloader:\n    images = images.to(device)\n    \n    with torch.no_grad():\n        outputs = model(images)\n        \n    _, preds = torch.max(outputs, 1)\n    \n    preds = preds.cpu()\n    predictions.extend(preds.tolist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:56:00.185401Z","iopub.execute_input":"2026-02-28T12:56:00.185684Z","iopub.status.idle":"2026-02-28T12:56:07.315733Z","shell.execute_reply.started":"2026-02-28T12:56:00.185661Z","shell.execute_reply":"2026-02-28T12:56:07.314947Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Формирование файла submission","metadata":{}},{"cell_type":"code","source":"predictions_df = pd.DataFrame({'ImageId': range(1, len(predictions) + 1), 'Label': predictions})\n\npredictions_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:56:12.582498Z","iopub.execute_input":"2026-02-28T12:56:12.583046Z","iopub.status.idle":"2026-02-28T12:56:12.591239Z","shell.execute_reply.started":"2026-02-28T12:56:12.583019Z","shell.execute_reply":"2026-02-28T12:56:12.590662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions_df.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T12:56:16.100537Z","iopub.execute_input":"2026-02-28T12:56:16.101693Z","iopub.status.idle":"2026-02-28T12:56:16.106479Z","shell.execute_reply.started":"2026-02-28T12:56:16.101655Z","shell.execute_reply":"2026-02-28T12:56:16.10594Z"}},"outputs":[],"execution_count":null}]}