{"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":[{"sourceId":30201,"databundleVersionId":2750748,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":730845,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":556736,"modelId":569301}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\nimport numpy as np\nimport pandas as pd\nimport random\nimport colorsys\nfrom tqdm.auto import tqdm\nimport json\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport cv2\nimport albumentations as A\n\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimport torch\nimport torchvision\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.transforms import functional as F\n\nfrom sklearn.model_selection import train_test_split\n\n# ====================== КОНФИГУРАЦИЯ ======================\nclass Config:\n    # Пути\n    ROOT = \"/kaggle/input/sartorius-cell-instance-segmentation\"\n    OUTPUT_DIR = \"/kaggle/working\"\n    MODEL_DIR = \"/kaggle/working/models\"\n    \n    # Путь к предобученным весам (измените на свой путь)\n    WEIGHTS_PATH = \"/kaggle/input/model-1/pytorch/default/1/models/model_final.pth\"  # Измените на свой путь\n    \n    # Параметры модели\n    NUM_CLASSES = 2\n    NUM_EPOCHS = 10\n    BATCH_SIZE = 2\n    LEARNING_RATE = 0.0001\n    \n    # Флаги (определяются автоматически)\n    TRAIN_MODE = True\n    USE_INTERNET = True\n    \n    @classmethod\n    def check_weights(cls):\n        \"\"\"Проверяет наличие весов по указанному пути\"\"\"\n        # Проверяем основной путь\n        if os.path.exists(cls.WEIGHTS_PATH):\n            print(f\"Найдены веса по указанному пути: {cls.WEIGHTS_PATH}\")\n            return True\n        \n        # Проверяем альтернативные пути\n        alt_paths = [\n            \"/kaggle/input/models-weights/models/model_final.pth\",\n            \"/kaggle/input/model-weights/model_final.pth\",\n            \"/kaggle/input/weights/model_final.pth\",\n            \"/kaggle/working/models/model_final.pth\"\n        ]\n        \n        for path in alt_paths:\n            if os.path.exists(path):\n                print(f\"Найдены веса по альтернативному пути: {path}\")\n                cls.WEIGHTS_PATH = path\n                return True\n        \n        print(\"Веса не найдены ни по одному из путей\")\n        return False\n    \n    @classmethod\n    def determine_mode(cls):\n        \"\"\"Определяет режим работы на основе наличия весов\"\"\"\n        if cls.check_weights():\n            print(\"Переключаемся в режим инференса.\")\n            cls.TRAIN_MODE = False\n            cls.USE_INTERNET = False\n        else:\n            print(\"Веса не найдены. Запускаем обучение.\")\n            cls.TRAIN_MODE = True\n            cls.USE_INTERNET = True\n\n# Инициализируем конфигурацию\nConfig.determine_mode()\n\n# ====================== УТИЛИТЫ ======================\ndef create_directories():\n    \"\"\"Создает необходимые директории\"\"\"\n    os.makedirs(Config.MODEL_DIR, exist_ok=True)\n    os.makedirs(os.path.join(Config.OUTPUT_DIR, \"checkpoints\"), exist_ok=True)\n\n# ====================== ФУНКЦИИ ИЗ ПРИМЕРА ======================\n\ndef decode_rle_mask(rle_mask, shape=(520, 704)):\n    \"\"\"Декодирует RLE маску в бинарное изображение\"\"\"\n    if pd.isna(rle_mask) or rle_mask.strip() == '':\n        return np.zeros(shape, dtype=np.uint8)\n    \n    rle_mask = rle_mask.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (rle_mask[0:][::2], rle_mask[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n\n    mask = np.zeros((shape[0] * shape[1]), dtype=np.uint8)\n    \n    for start, end in zip(starts, ends):\n        mask[start:end] = 1\n\n    mask = mask.reshape(shape[0], shape[1])\n    return np.uint8(mask)\n\ndef rle_encoding(x):\n    \"\"\"Кодирует бинарное изображение в RLE формат\"\"\"\n    dots = np.where(x.flatten() == 1)[0]\n    run_lengths = []\n    prev = -2\n    \n    for b in dots:\n        if (b > prev + 1): \n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n        \n    return ' '.join(map(str, run_lengths))\n\ndef remove_overlapping_pixels(mask, other_masks):\n    \"\"\"Удаляет перекрывающиеся пиксели из маски\"\"\"\n    for other_mask in other_masks:\n        if np.sum(np.logical_and(mask, other_mask)) > 0:\n            mask[np.logical_and(mask, other_mask)] = 0\n    return mask\n\n# Функция для преобразования тензора в изображение для отображения\ndef tensor_to_image(tensor, normalize=True):\n    \"\"\"\n    Преобразует тензор в изображение для отображения.\n    \n    Args:\n        tensor: Тензор в формате (C, H, W)\n        normalize: Если True, денормализует изображение\n        \n    Returns:\n        numpy array в формате (H, W, C) для отображения\n    \"\"\"\n    # Клонируем тензор, чтобы не изменять оригинал\n    tensor = tensor.clone().detach().cpu()\n    \n    if normalize:\n        # Денормализуем (для тестовых данных, которые нормализованы)\n        mean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\n        std = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n        tensor = tensor * std + mean\n    \n    # Приводим к диапазону [0, 1]\n    tensor = torch.clamp(tensor, 0, 1)\n    \n    # Преобразуем в numpy и меняем порядок осей\n    img = tensor.numpy().transpose(1, 2, 0)\n    \n    return img\n\n# Функция для отображения изображения без нормализации\ndef tensor_to_display(tensor):\n    \"\"\"\n    Преобразует тензор в изображение для отображения без нормализации.\n    Используется для обучающих данных, которые не нормализованы.\n    \"\"\"\n    # Клонируем тензор, чтобы не изменять оригинал\n    tensor = tensor.clone().detach().cpu()\n    \n    # Просто приводим к диапазону [0, 1] (изображения уже в диапазоне 0-255, делим на 255)\n    tensor = torch.clamp(tensor / 255.0, 0, 1)\n    \n    # Преобразуем в numpy и меняем порядок осей\n    img = tensor.numpy().transpose(1, 2, 0)\n    \n    return img\n\n# ====================== ДАТАСЕТ (точная копия из примера) ======================\nclass CellSegData(torch.utils.data.Dataset):\n    def __init__(self, root, df, split='train', aug=None, cls_map=None):\n        self.augmentations = aug\n        self.cls_map = cls_map\n        \n        train, test = train_test_split(df['id'].unique(), train_size=0.9)\n        \n        if split == 'train':\n            self.dataset = train\n        else:\n            self.dataset = test\n        \n        self.dict_df = {img_id: df[df['id']==img_id] for img_id in tqdm(self.dataset)}\n        self.root = root\n        \n    def __len__(self):\n        return len(self.dataset)\n    \n    def __getitem__(self, index):\n        img_id = self.dataset[index]\n        image = cv2.imread(os.path.join(self.root, img_id+'.png'))\n        \n        info = self.dict_df[img_id]\n        n_objects = len(info['annotation'])\n        \n        labels = info[\"cell_type\"].apply(lambda x: self.cls_map[x]).values\n        rles = info[\"annotation\"].values\n        \n        masks = []\n        for mask in rles:\n            decoded_mask = decode_rle_mask(rle_mask=mask, shape=image.shape[:2])\n            masks.append(decoded_mask)\n            \n        bboxes = []\n        for mask in masks:\n            pos = np.nonzero(mask)\n            xmin = np.min(pos[1])\n            xmax = np.max(pos[1])\n            ymin = np.min(pos[0])\n            ymax = np.max(pos[0])\n            bboxes.append([xmin, ymin, xmax, ymax])\n        \n        bboxes = np.asarray(bboxes)\n        \n        if self.augmentations is not None:\n            augmented = self.augmentations(\n                image=image, \n                masks=masks, \n                bboxes=bboxes,\n                labels=labels\n            )\n            image = augmented['image']\n            masks = augmented['masks']\n            bboxes = augmented['bboxes']\n            bboxes = np.stack(bboxes).astype(int)\n        \n        masks = np.asarray(masks)\n        bboxes = torch.as_tensor(bboxes, dtype=torch.int64)\n\n        is_bad_labels = False\n        degenerate_boxes = bboxes[:, 2:] <= bboxes[:, :2]\n        \n        if degenerate_boxes.any():\n            is_bad_labels = True\n            bb_idxs = torch.where(degenerate_boxes.any(dim=1))[0].numpy()\n\n        labels = torch.as_tensor([1 for i in range(len(bboxes))], dtype=torch.int64)\n        masks = torch.as_tensor(masks, dtype=torch.uint8)\n\n        if is_bad_labels:\n            bboxes = bboxes[[i for i in range(len(bboxes)) if i not in bb_idxs]]\n            labels = labels[[i for i in range(len(bboxes)) if i not in bb_idxs]]\n            masks = masks[[i for i in range(len(bboxes)) if i not in bb_idxs]]\n\n        image_id = torch.tensor([index])\n        area = (bboxes[:, 3] - bboxes[:, 1]) * (bboxes[:, 2] - bboxes[:, 0])\n        iscrowd = torch.zeros((n_objects,), dtype=torch.int64)\n\n        target = {\n            'boxes': bboxes,\n            'labels': labels,\n            'masks': masks,\n            'image_id': image_id,\n            'area': area,\n            'iscrowd': iscrowd\n        }\n        \n        image = image.transpose((2, 0, 1))\n        return torch.Tensor(image), target\n\n# ====================== МЕТРИКИ И УТИЛИТЫ ИЗ ПРИМЕРА ======================\nimport datetime\nimport errno\nimport time\nfrom collections import defaultdict, deque\n\nclass SmoothedValue:\n    def __init__(self, window_size=20, fmt=None):\n        if fmt is None:\n            fmt = \"{median:.4f} ({global_avg:.4f})\"\n        self.deque = deque(maxlen=window_size)\n        self.total = 0.0\n        self.count = 0\n        self.fmt = fmt\n\n    def update(self, value, n=1):\n        self.deque.append(value)\n        self.count += n\n        self.total += value * n\n\n    @property\n    def median(self):\n        d = torch.tensor(list(self.deque))\n        return d.median().item()\n\n    @property\n    def avg(self):\n        d = torch.tensor(list(self.deque), dtype=torch.float32)\n        return d.mean().item()\n\n    @property\n    def global_avg(self):\n        return self.total / self.count\n\n    @property\n    def max(self):\n        return max(self.deque)\n\n    @property\n    def value(self):\n        return self.deque[-1]\n\n    def __str__(self):\n        return self.fmt.format(\n            median=self.median, avg=self.avg, global_avg=self.global_avg, max=self.max, value=self.value\n        )\n\nclass MetricLogger:\n    def __init__(self, delimiter=\"\\t\"):\n        self.meters = defaultdict(SmoothedValue)\n        self.delimiter = delimiter\n\n    def add_meter(self, name, meter):\n        \"\"\"Добавить новый метр в логирование\"\"\"\n        self.meters[name] = meter\n\n    def update(self, **kwargs):\n        for k, v in kwargs.items():\n            if isinstance(v, torch.Tensor):\n                v = v.item()\n            assert isinstance(v, (float, int))\n            if k not in self.meters:\n                self.meters[k] = SmoothedValue()\n            self.meters[k].update(v)\n\n    def __getattr__(self, attr):\n        if attr in self.meters:\n            return self.meters[attr]\n        if attr in self.__dict__:\n            return self.__dict__[attr]\n        raise AttributeError(f\"'{type(self).__name__}' object has no attribute '{attr}'\")\n\n    def __str__(self):\n        loss_str = []\n        for name, meter in self.meters.items():\n            loss_str.append(f\"{name}: {str(meter)}\")\n        return self.delimiter.join(loss_str)\n\n    def log_every(self, iterable, print_freq, header=None):\n        i = 0\n        if not header:\n            header = \"\"\n        start_time = time.time()\n        end = time.time()\n        iter_time = SmoothedValue(fmt=\"{avg:.4f}\")\n        data_time = SmoothedValue(fmt=\"{avg:.4f}\")\n        space_fmt = \":\" + str(len(str(len(iterable)))) + \"d\"\n        \n        log_msg = self.delimiter.join(\n            [\n                header,\n                \"[{0\" + space_fmt + \"}/{1}]\",\n                \"eta: {eta}\",\n                \"{meters}\",\n                \"time: {time}\",\n                \"data: {data}\",\n                \"max mem: {memory:.0f}\",\n            ]\n        )\n        \n        MB = 1024.0 * 1024.0\n        for obj in iterable:\n            data_time.update(time.time() - end)\n            yield obj\n            iter_time.update(time.time() - end)\n            if i % print_freq == 0 or i == len(iterable) - 1:\n                eta_seconds = iter_time.global_avg * (len(iterable) - i)\n                eta_string = str(datetime.timedelta(seconds=int(eta_seconds)))\n                print(\n                    log_msg.format(\n                        i,\n                        len(iterable),\n                        eta=eta_string,\n                        meters=str(self),\n                        time=str(iter_time),\n                        data=str(data_time),\n                        memory=torch.cuda.max_memory_allocated() / MB,\n                    )\n                )\n            i += 1\n            end = time.time()\n        total_time = time.time() - start_time\n        total_time_str = str(datetime.timedelta(seconds=int(total_time)))\n        print(f\"{header} Total time: {total_time_str} ({total_time / len(iterable):.4f} s / it)\")\n\ndef collate_fn(batch):\n    return tuple(zip(*batch))\n\ndef save_on_master(*args, **kwargs):\n    torch.save(*args, **kwargs)\n\n# ====================== МОДЕЛЬ ======================\ndef create_model():\n    \"\"\"Создает и загружает модель\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"СОЗДАНИЕ МОДЕЛИ\")\n    print(\"=\"*60)\n    \n    # Создаем модель\n    model = torchvision.models.detection.maskrcnn_resnet50_fpn(\n        pretrained=False, \n        box_detections_per_img=600,\n        pretrained_backbone=False\n    )\n    \n    # Модифицируем головы\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, Config.NUM_CLASSES)\n    \n    in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\n    hidden_layer = 256\n    model.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, Config.NUM_CLASSES)\n    \n    # Загрузка весов в зависимости от режима\n    if not Config.TRAIN_MODE:\n        # Режим инференса - загружаем сохраненные веса\n        print(f\"Загружаем веса из: {Config.WEIGHTS_PATH}\")\n        try:\n            checkpoint = torch.load(Config.WEIGHTS_PATH, map_location='cpu')\n            \n            if 'model' in checkpoint:\n                model.load_state_dict(checkpoint['model'])\n                print(\"Веса успешно загружены (формат с ключом 'model')\")\n            elif 'state_dict' in checkpoint:\n                model.load_state_dict(checkpoint['state_dict'])\n                print(\"Веса успешно загружены (формат с ключом 'state_dict')\")\n            else:\n                # Пробуем загрузить напрямую\n                model.load_state_dict(checkpoint)\n                print(\"Веса успешно загружены (прямой формат)\")\n            \n            print(f\"Модель загружена для инференса\")\n            \n        except Exception as e:\n            print(f\"Ошибка при загрузке весов: {e}\")\n            print(\"Создаем модель с нуля для инференса\")\n    \n    elif Config.USE_INTERNET and Config.TRAIN_MODE:\n        # Режим обучения с загрузкой из интернета\n        print(\"Загружаем предобученную модель из интернета...\")\n        try:\n            # Создаем новую модель с предобученными весами\n            model = torchvision.models.detection.maskrcnn_resnet50_fpn(\n                pretrained=True,\n                box_detections_per_img=600\n            )\n            \n            # Модифицируем головы для нашего числа классов\n            in_features = model.roi_heads.box_predictor.cls_score.in_features\n            model.roi_heads.box_predictor = FastRCNNPredictor(in_features, Config.NUM_CLASSES)\n            \n            in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\n            hidden_layer = 256\n            model.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, Config.NUM_CLASSES)\n            \n            print(\"Модель успешно загружена из интернета\")\n            \n        except Exception as e:\n            print(f\"Ошибка загрузки из интернета: {e}\")\n            print(\"Продолжаем с текущей моделью (созданной с нуля)\")\n    \n    else:\n        # Режим обучения с нуля\n        print(\"Создаем модель с нуля для обучения\")\n    \n    return model\n\n# ====================== ФУНКЦИЯ АНАЛИЗА ======================\ndef analyze_train_sample(model, ds_train, sample_index, device):\n    \"\"\"Анализирует сэмпл из обучающего набора\"\"\"\n    img, targets = ds_train[sample_index]\n    \n    # Преобразуем тензор в изображение для отображения (без нормализации!)\n    img_display = tensor_to_display(img)\n    \n    # Получаем размер изображения\n    height, width = img_display.shape[0], img_display.shape[1]\n    \n    plt.figure(figsize=(12, 4))\n    \n    # Отображаем исходное изображение\n    plt.subplot(1, 3, 1)\n    plt.imshow(img_display)\n    plt.title(\"Original Image\")\n    plt.axis('off')\n\n    # Отображаем ground truth маски\n    plt.subplot(1, 3, 2)\n    plt.imshow(img_display)\n    \n    masks = np.zeros((height, width))\n    if len(targets['masks']) > 0:\n        for mask in targets['masks']:\n            masks = np.logical_or(masks, mask.numpy())\n    \n    plt.imshow(masks, alpha=0.3, cmap='jet')\n    plt.title(f\"Ground Truth: {len(targets['masks'])} masks\")\n    plt.axis('off')\n\n    # Отображаем предсказания модели\n    model.eval()\n    with torch.no_grad():\n        preds = model([img.to(device)])[0]\n\n    plt.subplot(1, 3, 3)\n    plt.imshow(img_display)\n    \n    all_preds_masks = np.zeros((height, width))\n    if len(preds['masks']) > 0:\n        for mask in preds['masks'].cpu().detach().numpy():\n            all_preds_masks = np.logical_or(all_preds_masks, mask[0] > 0.5)\n    \n    plt.imshow(all_preds_masks, alpha=0.3, cmap='jet')\n    plt.title(f\"Predictions: {len(preds['masks'])} masks\")\n    plt.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n    \n    # Выводим статистику\n    print(f\"Image shape: {img.shape}\")\n    print(f\"Image min/max: {img.min():.3f}/{img.max():.3f}\")\n    print(f\"Display image min/max: {img_display.min():.3f}/{img_display.max():.3f}\")\n    print(f\"Ground Truth: {len(targets['masks'])} объектов\")\n    print(f\"Predictions: {len(preds['masks'])} обнаружено, {np.sum(all_preds_masks > 0)} пикселей маски\")\n\n# ====================== ТЕСТОВЫЙ ДАТАСЕТ (из примера) ======================\nclass CellTestDataset(torch.utils.data.Dataset):\n    def __init__(self, image_dir, transforms=None):\n        self.transforms = transforms\n        self.image_dir = image_dir\n        self.image_ids = [f[:-4] for f in os.listdir(self.image_dir)]\n        \n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        image_path = os.path.join(self.image_dir, image_id + \".png\")\n        image = Image.open(image_path).convert(\"RGB\")\n\n        if self.transforms is not None:\n            image, _ = self.transforms(image=image, target=None)\n            \n        return {'image': image, 'image_id': image_id}\n\n    def __len__(self):\n        return len(self.image_ids)\n\n# Transform для теста (из примера)\nclass Compose:\n    def __init__(self, transforms):\n        self.transforms = transforms\n\n    def __call__(self, image, target):\n        for t in self.transforms:\n            image, target = t(image, target)\n        return image, target\n\nclass Normalize:\n    def __call__(self, image, target):\n        image = F.normalize(image, (0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n        return image, target\n    \nclass ToTensor:\n    def __call__(self, image, target):\n        image = F.to_tensor(image)\n        return image, target\n    \ndef get_transform(train):\n    transforms = [ToTensor()]\n    if True:\n        transforms.append(Normalize())\n    return Compose(transforms)\n\n# ====================== ОБУЧЕНИЕ ======================\ndef train_model(model, data_loader, data_loader_test, dataset, device, num_epochs=10):\n    \"\"\"Обучение модели\"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"НАЧАЛО ОБУЧЕНИЯ\")\n    print(\"=\"*60)\n    \n    params = [p for p in model.parameters() if p.requires_grad]\n    optimizer = torch.optim.AdamW(params, lr=0.0001)\n    lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)\n\n    output_dir = 'weights'\n    os.makedirs(output_dir, exist_ok=True)\n    \n    for epoch in range(num_epochs):\n        # train for one epoch, printing every 10 iterations\n        model.train()\n            \n        metric_logger = MetricLogger(delimiter=\"  \")\n        metric_logger.add_meter(\"lr\", SmoothedValue(window_size=1, fmt=\"{value:.6f}\"))\n        header = f\"Epoch: [{epoch+1}/{num_epochs}]\"\n\n        lr_scheduler = None\n            \n        if epoch == 0:\n            warmup_factor = 1.0 / 1000\n            warmup_iters = min(1000, len(data_loader) - 1)\n\n            lr_scheduler = torch.optim.lr_scheduler.LinearLR(\n                optimizer, start_factor=warmup_factor, total_iters=warmup_iters\n            )\n\n        for images, targets in metric_logger.log_every(data_loader, 10, header):\n            images = list(image.to(device) for image in images)\n            targets = [{k: v.to(device) for k, v in t.items()} for t in targets]\n\n            loss_dict = model(images, targets)\n            losses = sum(loss for loss in loss_dict.values())\n\n            optimizer.zero_grad()\n            losses.backward()\n            optimizer.step()\n\n            if lr_scheduler is not None:\n                lr_scheduler.step()\n\n            metric_logger.update(loss=losses, **loss_dict)\n            metric_logger.update(lr=optimizer.param_groups[0][\"lr\"])\n\n        # Проанализируем один сэмпл\n        if epoch % 2 == 0:  # Анализируем каждые 2 эпохи\n            analyze_train_sample(model, dataset, 20, device)\n        \n        # Сохраняем чекпоинт\n        if output_dir:\n            checkpoint = {\n                \"model\": model.state_dict(),\n                \"optimizer\": optimizer.state_dict(),\n                \"epoch\": epoch,\n                \"config\": {\n                    'num_classes': Config.NUM_CLASSES,\n                    'box_detections_per_img': 600\n                }\n            }\n            save_on_master(checkpoint, os.path.join(output_dir, f\"model_{epoch}.pth\"))\n            save_on_master(checkpoint, os.path.join(output_dir, \"checkpoint.pth\"))\n            \n            # Также сохраняем в основную директорию\n            torch.save(model.state_dict(), os.path.join(Config.MODEL_DIR, f\"model_epoch_{epoch}.pth\"))\n    \n    return model\n\n# ====================== СОЗДАНИЕ SUBMISSION ======================\ndef create_submission_final(model, test_dataset, device):\n    \"\"\"Создает файл submission для Kaggle\"\"\"\n    model.eval()\n    submission = []\n    \n    print(f\"Обрабатываем {len(test_dataset)} тестовых изображений...\")\n    \n    for idx, sample in enumerate(test_dataset):\n        img = sample['image']\n        image_id = sample['image_id']\n        \n        with torch.no_grad():\n            result = model([img.to(device)])[0]\n        \n        previous_masks = []\n        masks_per_image = []\n        \n        for i, mask in enumerate(result[\"masks\"]):\n            # Filter-out low-scoring results\n            score = result[\"scores\"][i].cpu().item()\n            if score < 0.3:\n                continue\n            \n            mask = mask.cpu().numpy()\n            # Keep only highly likely pixels\n            binary_mask = mask > 0.5\n            binary_mask = remove_overlapping_pixels(binary_mask, previous_masks)\n            previous_masks.append(binary_mask)\n            rle = rle_encoding(binary_mask)\n            masks_per_image.append(rle)\n        \n        # Объединяем все маски для этого изображения\n        for rle in masks_per_image:\n            submission.append((image_id, rle))\n        \n        # Визуализация (только для первых 3 изображений)\n        if idx < 3:\n            # Для тестовых данных используем денормализацию\n            img_display = tensor_to_image(img.clone(), normalize=True)\n            \n            plt.figure(figsize=(12,6))\n            ax1 = plt.subplot(121)\n            ax1.imshow(img_display)\n            ax1.set_title(f\"Original: {image_id}\")\n            ax1.axis('off')\n            \n            all_preds_masks = np.zeros((520, 704))\n            \n            for mask in result['masks'].cpu().detach().numpy():\n                all_preds_masks = np.logical_or(all_preds_masks, mask[0] > 0.5)\n                \n            ax2 = plt.subplot(122)\n            ax2.imshow(img_display)\n            ax2.imshow(all_preds_masks, alpha=0.3, cmap='jet')\n            ax2.set_title(f\"Predictions: {len(masks_per_image)} masks\")\n            ax2.axis('off')\n            plt.tight_layout()\n            plt.show()\n        \n        # Прогресс\n        if (idx + 1) % 10 == 0:\n            print(f\"Обработано {idx + 1}/{len(test_dataset)} изображений\")\n    \n    df_sub = pd.DataFrame(submission, columns=['id', 'predicted'])\n    \n    # Сохраняем submission\n    submission_path = os.path.join(Config.OUTPUT_DIR, \"submission.csv\")\n    df_sub.to_csv(submission_path, index=False)\n    \n    print(f\"\\nSubmission создан: {submission_path}\")\n    print(f\"Количество строк: {len(df_sub)}\")\n    print(f\"Уникальных изображений: {df_sub['id'].nunique()}\")\n    \n    # Статистика\n    non_empty = df_sub[df_sub['predicted'].str.strip() != '']\n    print(f\"Непустых предсказаний: {len(non_empty)}\")\n    \n    print(f\"\\nПервые 5 строк submission:\")\n    print(df_sub.head())\n    \n    return df_sub\n\n# ====================== ОСНОВНОЙ БЛОК ======================\ndef main():\n    print(\"=\" * 60)\n    print(f\"РЕЖИМ РАБОТЫ: {'ОБУЧЕНИЕ' if Config.TRAIN_MODE else 'ИНФЕРЕНС'}\")\n    print(\"=\" * 60)\n    \n    # Создаем директории\n    create_directories()\n    \n    # Устройство\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Устройство: {device}\")\n    \n    if Config.TRAIN_MODE:\n        # ==================== РЕЖИМ ОБУЧЕНИЯ ====================\n        print(\"\\n\" + \"=\"*60)\n        print(\"НАЧИНАЕМ ОБУЧЕНИЕ\")\n        print(\"=\"*60)\n        \n        # Загружаем данные\n        print(\"Загружаем данные...\")\n        train_df = pd.read_csv(os.path.join(Config.ROOT, \"train.csv\"))\n        print(f\"Размер train данных: {train_df.shape}\")\n        \n        # Создаем маппинг классов\n        cls_map = {value: idx for idx, value in enumerate(train_df[\"cell_type\"].unique())}\n        print(f\"Маппинг классов: {cls_map}\")\n        \n        # Аугментации\n        train_augmentations = A.Compose([\n            A.HorizontalFlip(),\n            A.VerticalFlip(),\n            A.RandomRotate90()\n        ], bbox_params={\n            \"format\": \"pascal_voc\",\n            \"min_area\": 0,\n            \"min_visibility\": 0,\n            'label_fields': ['labels']\n        })\n        \n        # Создаем датасеты\n        print(\"Создаем датасеты...\")\n        dataset = CellSegData(\n            os.path.join(Config.ROOT, \"train\"),\n            train_df, \n            'train', \n            train_augmentations, \n            cls_map\n        )\n        \n        dataset_test = CellSegData(\n            os.path.join(Config.ROOT, \"train\"),\n            train_df, \n            'test', \n            train_augmentations, \n            cls_map\n        )\n        \n        print(f\"Размер train датасета: {len(dataset)}\")\n        print(f\"Размер test датасета: {len(dataset_test)}\")\n        \n        # DataLoader\n        data_loader = torch.utils.data.DataLoader(\n            dataset, \n            batch_size=2, \n            shuffle=True, \n            num_workers=0,\n            collate_fn=collate_fn\n        )\n\n        data_loader_test = torch.utils.data.DataLoader(\n            dataset_test, \n            batch_size=1, \n            shuffle=False, \n            num_workers=0, \n            collate_fn=collate_fn\n        )\n        \n        # Создаем модель\n        model = create_model()\n        model.to(device)\n        \n        # Тренируем модель\n        print(f\"\\nНачинаем обучение на {Config.NUM_EPOCHS} эпох...\")\n        model = train_model(model, data_loader, data_loader_test, dataset, device, Config.NUM_EPOCHS)\n        \n        # Сохраняем финальную модель\n        final_model_path = os.path.join(Config.MODEL_DIR, \"model_final.pth\")\n        checkpoint = {\n            \"model\": model.state_dict(),\n            \"epoch\": Config.NUM_EPOCHS - 1,\n            \"config\": {\n                'num_classes': Config.NUM_CLASSES,\n                'box_detections_per_img': 600\n            }\n        }\n        torch.save(checkpoint, final_model_path)\n        print(f\"Финальная модель сохранена в {final_model_path}\")\n        \n        print(\"\\n\" + \"=\"*60)\n        print(\"ОБУЧЕНИЕ ЗАВЕРШЕНО!\")\n        print(\"=\"*60)\n        \n    else:\n        # ==================== РЕЖИМ ИНФЕРЕНСА ====================\n        print(\"\\n\" + \"=\"*60)\n        print(\"НАЧИНАЕМ ИНФЕРЕНС\")\n        print(\"=\"*60)\n        \n        # Создаем модель с загруженными весами\n        model = create_model()\n        model.to(device)\n        model.eval()\n        \n        # Создаем тестовый датасет\n        print(\"Создаем тестовый датасет...\")\n        ds_test = CellTestDataset(\n            f'{Config.ROOT}/test', \n            transforms=get_transform(train=False)\n        )\n        print(f\"Найдено тестовых изображений: {len(ds_test)}\")\n        \n        # Создаем submission\n        print(\"\\nСоздаем submission...\")\n        df_sub = create_submission_final(model, ds_test, device)\n        \n        print(\"\\n\" + \"=\"*60)\n        print(\"ИНФЕРЕНС ЗАВЕРШЕН!\")\n        print(\"=\"*60)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"ВЫПОЛНЕНИЕ ЗАВЕРШЕНО!\")\n    print(\"=\"*60)\n\n# Запуск основной функции\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T16:35:33.739664Z","iopub.execute_input":"2026-01-25T16:35:33.739947Z","iopub.status.idle":"2026-01-25T16:36:21.199864Z","shell.execute_reply.started":"2026-01-25T16:35:33.739921Z","shell.execute_reply":"2026-01-25T16:36:21.199130Z"}},"outputs":[],"execution_count":null}]}