{"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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":30201,"databundleVersionId":2750748}],"dockerImageVersionId":31329,"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        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":"2026-03-30T21:51:57.682854Z","iopub.execute_input":"2026-03-30T21:51:57.683283Z","iopub.status.idle":"2026-03-30T21:52:01.367520Z","shell.execute_reply.started":"2026-03-30T21:51:57.683241Z","shell.execute_reply":"2026-03-30T21:52:01.366769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\n\nimport cv2\nimport colorsys\nimport albumentations as A\nfrom collections import defaultdict, deque\nimport datetime\nimport time\n\n\nimport torch\nimport torchvision\nfrom torchvision.models.detection import maskrcnn_resnet50_fpn\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\n\nfrom sklearn.model_selection import train_test_split\n\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:11:39.011323Z","iopub.execute_input":"2026-03-30T22:11:39.012046Z","iopub.status.idle":"2026-03-30T22:11:39.017450Z","shell.execute_reply.started":"2026-03-30T22:11:39.012002Z","shell.execute_reply":"2026-03-30T22:11:39.016693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def decode_rle_mask(rle_mask, shape=(520, 704)):\n    \"\"\"\n    Декодирование RLE маски в 2D массив\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    for start, end in zip(starts, ends):\n        mask[start:end] = 1\n\n    mask = mask.reshape(shape[0], shape[1])\n    mask = np.uint8(mask)\n    return mask\n\n\ndef encode_rle_mask(mask):\n    \"\"\"\n    Кодирование маски в RLE формат\n    \"\"\"\n    pixels = mask.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    rle = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    rle[1::2] -= rle[::2]\n    return rle.tolist()\n\n\ndef get_bboxes_from_mask(masks):\n    \"\"\"\n    Получение bounding boxes из масок\n    \"\"\"\n    coco_boxes = []\n    for mask in masks:\n        pos = np.nonzero(mask)\n        if len(pos[0]) > 0 and len(pos[1]) > 0:\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            coco_boxes.append([xmin, ymin, xmax, ymax])\n    return np.asarray(coco_boxes) if coco_boxes else np.array([]).reshape(0, 4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T21:53:58.482453Z","iopub.execute_input":"2026-03-30T21:53:58.482753Z","iopub.status.idle":"2026-03-30T21:53:58.491187Z","shell.execute_reply.started":"2026-03-30T21:53:58.482726Z","shell.execute_reply":"2026-03-30T21:53:58.490518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def random_colors(N, bright=True):\n    \"\"\"\n    Генерация случайных цветов для визуализации\n    \"\"\"\n    brightness = 1.0 if bright else 0.7\n    hsv = [(i / N, 1, brightness) for i in range(N)]\n    colors = list(map(lambda c: colorsys.hsv_to_rgb(*c), hsv))\n    random.shuffle(colors)\n    return colors\n\n\ndef apply_mask(image, mask, color, alpha=0.5):\n    \"\"\"\n    Применение маски к изображению\n    \"\"\"\n    for c in range(3):\n        image[:, :, c] = np.where(mask == 1,\n                                  image[:, :, c] * (1 - alpha) + alpha * color[c] * 255,\n                                  image[:, :, c])\n    return image\n\n\ndef plot_image_annotations(image, masks, bboxes, labels, aug=None):\n    \"\"\"\n    Визуализация изображения с аннотациями\n    \"\"\"\n    image = image.copy()\n    \n    colors = {unique_lbl: random_colors(1, True)[0] for unique_lbl in np.unique(labels)}\n    \n    if aug is not None:\n        augmented = aug(image=image, masks=masks, bboxes=bboxes, labels=labels)\n        image = augmented['image']\n        masks = augmented['masks']\n        bboxes = augmented['bboxes']\n    \n    if len(bboxes) > 0:\n        bboxes = np.stack(bboxes).astype(int)\n        \n        for idx, box in enumerate(bboxes):\n            color = tuple([int(value * 255) for value in colors[labels[idx]]])\n            image = cv2.rectangle(image, (box[2], box[3]), (box[0], box[1]), color=color, thickness=2)\n    \n    for idx, mask in enumerate(masks):\n        color = colors[labels[idx]]\n        image = apply_mask(image, mask, color)\n    \n    plt.figure(figsize=(15, 15))\n    plt.imshow(image)\n    plt.axis('off')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T21:54:24.673464Z","iopub.execute_input":"2026-03-30T21:54:24.673765Z","iopub.status.idle":"2026-03-30T21:54:24.682329Z","shell.execute_reply.started":"2026-03-30T21:54:24.673739Z","shell.execute_reply":"2026-03-30T21:54:24.681549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Загрузка данных\ntrain_df = pd.read_csv('/kaggle/input/competitions/sartorius-cell-instance-segmentation/train.csv')\nprint(f\"Training data shape: {train_df.shape}\")\nprint(f\"Number of unique images: {train_df['id'].nunique()}\")\nprint(f\"Cell types: {train_df['cell_type'].unique()}\")\nprint(f\"Value counts per cell type:\\n{train_df['cell_type'].value_counts()}\")\n\n# Создание словаря классов\ncls_map = {value: idx for idx, value in enumerate(train_df['cell_type'].unique())}\ncls_map_reversed = {v: k for k, v in cls_map.items()}\nprint(f\"\\nClass mapping: {cls_map}\")\n\n# Статистика по аннотациям\ntrain_df['annotation_length'] = train_df['annotation'].apply(lambda x: len(x.split()))\nprint(f\"\\nAnnotation stats:\")\nprint(train_df['annotation_length'].describe())\n\n# Просмотр первых строк\nprint(\"\\nFirst few rows:\")\nprint(train_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:02:19.701969Z","iopub.execute_input":"2026-03-30T22:02:19.702579Z","iopub.status.idle":"2026-03-30T22:02:20.187161Z","shell.execute_reply.started":"2026-03-30T22:02:19.702550Z","shell.execute_reply":"2026-03-30T22:02:20.186420Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Путь к данным\nDATA_PATH = '/kaggle/input/competitions/sartorius-cell-instance-segmentation'\n\n# Покажем несколько примеров изображений с масками\nsample_ids = train_df['id'].unique()[:3]\n\nfor img_id in sample_ids:\n    # Загрузка изображения\n    img_path = os.path.join(DATA_PATH, 'train', f'{img_id}.png')\n    img = cv2.imread(img_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    \n    # Получение аннотаций для этого изображения\n    img_data = train_df[train_df['id'] == img_id]\n    labels = img_data['cell_type'].values\n    rles = img_data['annotation'].values\n    \n    # Декодирование масок\n    masks = []\n    for rle in rles:\n        mask = decode_rle_mask(rle_mask=rle, shape=img.shape)\n        masks.append(mask)\n    \n    # Получение bboxes\n    bboxes = get_bboxes_from_mask(masks)\n    \n    # Визуализация\n    print(f\"\\nImage ID: {img_id}\")\n    print(f\"Number of cells: {len(masks)}\")\n    print(f\"Cell types: {labels}\")\n    \n    plot_image_annotations(img, masks, bboxes, labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:06:11.337142Z","iopub.execute_input":"2026-03-30T22:06:11.337869Z","iopub.status.idle":"2026-03-30T22:06:15.877188Z","shell.execute_reply.started":"2026-03-30T22:06:11.337834Z","shell.execute_reply":"2026-03-30T22:06:15.876480Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class 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_ids, test_ids = train_test_split(df['id'].unique(), train_size=0.9, random_state=42)\n        if split == 'train':\n            self.dataset = train_ids\n        else:\n            self.dataset = test_ids\n            \n        self.dict_df = {img_id: df[df['id'] == img_id] for img_id in tqdm(self.dataset, desc=f\"Loading {split} data\")}\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        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\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)\n            masks.append(decoded_mask)\n            \n        bboxes = get_bboxes_from_mask(masks)\n        \n        if self.augmentations is not None and len(masks) > 0:\n            augmented = self.augmentations(image=image, masks=masks, bboxes=bboxes, labels=labels)\n            image = augmented['image']\n            masks = augmented['masks']\n            bboxes = augmented['bboxes']\n            if len(bboxes) > 0:\n                bboxes = np.stack(bboxes).astype(int)\n        \n        # Преобразование в тензоры\n        if len(bboxes) > 0:\n            bboxes = torch.as_tensor(bboxes, dtype=torch.float32)\n        else:\n            bboxes = torch.zeros((0, 4), dtype=torch.float32)\n            \n        labels = torch.as_tensor(labels, dtype=torch.int64)\n        \n        if len(masks) > 0:\n            masks = np.asarray(masks)\n            masks = torch.as_tensor(masks, dtype=torch.uint8)\n        else:\n            masks = torch.zeros((0, image.shape[0], image.shape[1]), dtype=torch.uint8)\n\n        image_id = torch.tensor([index])\n        \n        if len(bboxes) > 0:\n            area = (bboxes[:, 3] - bboxes[:, 1]) * (bboxes[:, 2] - bboxes[:, 0])\n        else:\n            area = torch.zeros((0,), dtype=torch.float32)\n            \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.astype(np.float32) / 255.0\n        image = image.transpose((2, 0, 1))\n        return torch.FloatTensor(image), target","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:07:02.386491Z","iopub.execute_input":"2026-03-30T22:07:02.386808Z","iopub.status.idle":"2026-03-30T22:07:02.403217Z","shell.execute_reply.started":"2026-03-30T22:07:02.386780Z","shell.execute_reply":"2026-03-30T22:07:02.402472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestCellSegData(torch.utils.data.Dataset):\n    def __init__(self, root, df, aug=None):\n        self.augmentations = aug\n        self.df = df\n        self.root = root\n        self.ids = df['id'].unique()\n        \n    def __len__(self):\n        return len(self.ids)\n    \n    def __getitem__(self, index):\n        img_id = self.ids[index]\n        image = cv2.imread(os.path.join(self.root, img_id + '.png'))\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        if self.augmentations is not None:\n            augmented = self.augmentations(image=image)\n            image = augmented['image']\n        \n        image = image.astype(np.float32) / 255.0\n        image = image.transpose((2, 0, 1))\n        return torch.FloatTensor(image), img_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:07:24.673849Z","iopub.execute_input":"2026-03-30T22:07:24.674237Z","iopub.status.idle":"2026-03-30T22:07:24.680117Z","shell.execute_reply.started":"2026-03-30T22:07:24.674208Z","shell.execute_reply":"2026-03-30T22:07:24.679396Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class 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 value(self):\n        return self.deque[-1]\n\n    def __str__(self):\n        return self.fmt.format(\n            median=self.median,\n            avg=self.avg,\n            global_avg=self.global_avg,\n            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 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            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 add_meter(self, name, meter):\n        self.meters[name] = meter\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        if torch.cuda.is_available():\n            log_msg = self.delimiter.join([\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        else:\n            log_msg = self.delimiter.join([\n                header,\n                '[{0' + space_fmt + '}/{1}]',\n                'eta: {eta}',\n                '{meters}',\n                'time: {time}',\n                'data: {data}'\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                if torch.cuda.is_available():\n                    print(log_msg.format(\n                        i, len(iterable), eta=eta_string,\n                        meters=str(self),\n                        time=str(iter_time), data=str(data_time),\n                        memory=torch.cuda.max_memory_allocated() / MB))\n                else:\n                    print(log_msg.format(\n                        i, len(iterable), eta=eta_string,\n                        meters=str(self),\n                        time=str(iter_time), data=str(data_time)))\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\n\ndef collate_fn(batch):\n    return tuple(zip(*batch))\n\n\ndef reduce_dict(input_dict, average=True):\n    world_size = 1\n    if world_size < 2:\n        return input_dict\n    with torch.no_grad():\n        names = []\n        values = []\n        for k in sorted(input_dict.keys()):\n            names.append(k)\n            values.append(input_dict[k])\n        values = torch.stack(values, dim=0)\n        reduced_dict = {k: v for k, v in zip(names, values)}\n    return reduced_dict\n\n\ndef save_on_master(*args, **kwargs):\n    torch.save(*args, **kwargs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:08:00.779065Z","iopub.execute_input":"2026-03-30T22:08:00.779780Z","iopub.status.idle":"2026-03-30T22:08:00.794933Z","shell.execute_reply.started":"2026-03-30T22:08:00.779750Z","shell.execute_reply":"2026-03-30T22:08:00.794198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_submission(model, test_loader, device, threshold=0.5):\n    model.eval()\n    submission_data = []\n    \n    with torch.no_grad():\n        for images, img_ids in tqdm(test_loader, desc=\"Predicting\"):\n            images = list(image.to(device) for image in images)\n            outputs = model(images)\n            \n            for i, output in enumerate(outputs):\n                masks = output['masks'].cpu().numpy()\n                scores = output['scores'].cpu().numpy()\n                labels = output['labels'].cpu().numpy()\n                img_id = img_ids[i]\n                \n                # Фильтруем по уверенности\n                keep = scores >= threshold\n                masks = masks[keep]\n                labels = labels[keep]\n                \n                for mask, label in zip(masks, labels):\n                    binary_mask = (mask[0] > 0.5).astype(np.uint8)\n                    rle = encode_rle_mask(binary_mask)\n                    submission_data.append({\n                        'id': img_id,\n                        'annotation': ' '.join(map(str, rle)),\n                        'cell_type': cls_map_reversed[label]\n                    })\n    \n    submission_df = pd.DataFrame(submission_data)\n    submission_df.to_csv('submission.csv', index=False)\n    print(f\"Submission saved! Total predictions: {len(submission_df)}\")\n    return submission_df\n\n\ndef load_test_data():\n    \"\"\"Загрузка тестовых данных\"\"\"\n    DATA_PATH = '/kaggle/input/competitions/sartorius-cell-instance-segmentation'\n    test_path = os.path.join(DATA_PATH, 'test.csv')\n    \n    if os.path.exists(test_path):\n        test_df = pd.read_csv(test_path)\n        print(f\"Test data loaded. Found {len(test_df)} test images\")\n    else:\n        test_images = [f.replace('.png', '') for f in os.listdir(os.path.join(DATA_PATH, 'test')) \n                      if f.endswith('.png')]\n        test_df = pd.DataFrame({'id': test_images})\n        print(f\"Created test dataframe from images. Found {len(test_df)} test images\")\n    return test_df\n\n\ndef analyze_train_sample(model, ds_train, sample_index):\n    \"\"\"Анализ предсказаний на обучающем примере\"\"\"\n    model.eval()\n    device = next(model.parameters()).device\n    \n    img, targets = ds_train[sample_index]\n    img = img.unsqueeze(0).to(device)\n    \n    with torch.no_grad():\n        preds = model(img)[0]\n    \n    img_np = img.cpu().squeeze(0).numpy().transpose((1, 2, 0))\n    \n    plt.figure(figsize=(15, 5))\n    \n    plt.subplot(1, 3, 1)\n    plt.imshow(img_np)\n    plt.title(\"Original Image\")\n    plt.axis('off')\n    \n    plt.subplot(1, 3, 2)\n    gt_masks = np.zeros((img_np.shape[0], img_np.shape[1]))\n    for mask in targets['masks']:\n        gt_masks = np.logical_or(gt_masks, mask.numpy())\n    plt.imshow(img_np)\n    plt.imshow(gt_masks, alpha=0.3)\n    plt.title(\"Ground Truth\")\n    plt.axis('off')\n    \n    plt.subplot(1, 3, 3)\n    pred_masks = np.zeros((img_np.shape[0], img_np.shape[1]))\n    for mask in preds['masks'].cpu().numpy():\n        pred_masks = np.logical_or(pred_masks, mask[0] > 0.5)\n    plt.imshow(img_np)\n    plt.imshow(pred_masks, alpha=0.3)\n    plt.title(\"Predictions\")\n    plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:08:26.112675Z","iopub.execute_input":"2026-03-30T22:08:26.113218Z","iopub.status.idle":"2026-03-30T22:08:26.127777Z","shell.execute_reply.started":"2026-03-30T22:08:26.113188Z","shell.execute_reply":"2026-03-30T22:08:26.127076Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Путь к данным\nDATA_PATH = '/kaggle/input/competitions/sartorius-cell-instance-segmentation'\n\n# Аугментации для обучения\ntrain_augmentations = A.Compose([\n    A.Resize(640, 640),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.RandomBrightnessContrast(p=0.3),\n], bbox_params=A.BboxParams(format=\"pascal_voc\", min_area=0, min_visibility=0, label_fields=['labels']))\n\n# Аугментации для теста (только resize)\ntest_augmentations = A.Compose([\n    A.Resize(640, 640)\n])\n\n# Параметры обучения\nNUM_EPOCHS = 10  # Начните с 10 эпох, затем увеличьте\nBATCH_SIZE = 2\nNUM_WORKERS = 2\nLEARNING_RATE = 0.0001\n\n# Устройство для обучения\ndevice = torch.device('cuda' if torch.cuda.is_available() else torch.device('cpu'))\nprint(f\"Using device: {device}\")\n\n# Если CUDA недоступна, уменьшаем batch_size\nif not torch.cuda.is_available():\n    BATCH_SIZE = 1\n    NUM_WORKERS = 0\n    print(f\"CUDA not available, reducing batch_size to {BATCH_SIZE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:08:44.284865Z","iopub.execute_input":"2026-03-30T22:08:44.285169Z","iopub.status.idle":"2026-03-30T22:08:44.294849Z","shell.execute_reply.started":"2026-03-30T22:08:44.285142Z","shell.execute_reply":"2026-03-30T22:08:44.294331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Создание датасетов\ndataset = CellSegData(os.path.join(DATA_PATH, 'train'), \n                      train_df, 'train', train_augmentations, cls_map)\ndataset_val = CellSegData(os.path.join(DATA_PATH, 'train'), \n                          train_df, 'test', train_augmentations, cls_map)\n\nprint(f\"Train dataset size: {len(dataset)}\")\nprint(f\"Validation dataset size: {len(dataset_val)}\")\n\n# Создание DataLoaders\ndata_loader = torch.utils.data.DataLoader(\n    dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS,\n    collate_fn=collate_fn)\n\ndata_loader_val = torch.utils.data.DataLoader(\n    dataset_val, batch_size=1, shuffle=False, num_workers=NUM_WORKERS,\n    collate_fn=collate_fn)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:09:02.540269Z","iopub.execute_input":"2026-03-30T22:09:02.541002Z","iopub.status.idle":"2026-03-30T22:09:05.926783Z","shell.execute_reply.started":"2026-03-30T22:09:02.540973Z","shell.execute_reply":"2026-03-30T22:09:05.925996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Создание модели Mask R-CNN\nNUM_CLASSES = len(cls_map) + 1  # +1 для фона\nprint(f\"Number of classes: {NUM_CLASSES}\")\n\nmodel = maskrcnn_resnet50_fpn(pretrained=True, box_detections_per_img=600)\n\n# Замена головы для классификации\nin_features = model.roi_heads.box_predictor.cls_score.in_features\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES)\n\n# Замена головы для масок\nin_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\nhidden_layer = 256\nmodel.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, NUM_CLASSES)\n\nmodel.to(device)\n\nprint(f\"Model created with {NUM_CLASSES} classes\")\nprint(f\"Total parameters: {sum(p.numel() for p in model.parameters())}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:09:20.772714Z","iopub.execute_input":"2026-03-30T22:09:20.773309Z","iopub.status.idle":"2026-03-30T22:09:22.694407Z","shell.execute_reply.started":"2026-03-30T22:09:20.773268Z","shell.execute_reply":"2026-03-30T22:09:22.693655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Оптимизатор\nparams = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.AdamW(params, lr=LEARNING_RATE)\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\n\n# Папка для сохранения весов\noutput_dir = 'weights'\nos.makedirs(output_dir, exist_ok=True)\n\n# Обучение\nbest_loss = float('inf')\n\nfor epoch in range(NUM_EPOCHS):\n    model.train()\n    metric_logger = MetricLogger(delimiter=\" \")\n    metric_logger.add_meter(\"lr\", SmoothedValue(window_size=1, fmt=\"{value:.6f}\"))\n    header = f\"Epoch: [{epoch}]\"\n\n    for images, targets in metric_logger.log_every(data_loader, 20, 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        loss_dict_reduced = reduce_dict(loss_dict)\n        losses_reduced = sum(loss for loss in loss_dict_reduced.values())\n\n        optimizer.zero_grad()\n        losses.backward()\n        optimizer.step()\n\n        metric_logger.update(loss=losses_reduced, **loss_dict_reduced)\n        metric_logger.update(lr=optimizer.param_groups[0][\"lr\"])\n\n    lr_scheduler.step()\n    \n    # Сохранение лучшей модели\n    if losses_reduced < best_loss:\n        best_loss = losses_reduced\n        checkpoint = {\n            \"model\": model.state_dict(),\n            \"optimizer\": optimizer.state_dict(),\n            \"epoch\": epoch,\n            \"loss\": best_loss,\n        }\n        torch.save(checkpoint, os.path.join(output_dir, \"best_model.pth\"))\n        print(f\"Best model saved with loss: {best_loss:.4f}\")\n    \n    # Сохранение чекпоинта каждой эпохи\n    checkpoint = {\n        \"model\": model.state_dict(),\n        \"optimizer\": optimizer.state_dict(),\n        \"epoch\": epoch,\n    }\n    torch.save(checkpoint, os.path.join(output_dir, f\"model_epoch_{epoch}.pth\"))\n    torch.save(checkpoint, os.path.join(output_dir, \"checkpoint.pth\"))\n    \n    # Визуализация предсказаний после каждой эпохи\n    if epoch % 2 == 0:\n        analyze_train_sample(model, dataset, np.random.randint(0, len(dataset)))\n\nprint(\"Training completed!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:11:52.688127Z","iopub.execute_input":"2026-03-30T22:11:52.688793Z","iopub.status.idle":"2026-03-30T22:44:58.443842Z","shell.execute_reply.started":"2026-03-30T22:11:52.688763Z","shell.execute_reply":"2026-03-30T22:44:58.442757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Загрузка лучшей модели\ncheckpoint = torch.load(os.path.join(output_dir, \"best_model.pth\"), map_location=device)\nmodel.load_state_dict(checkpoint['model'])\nmodel.eval()\nprint(f\"Loaded best model from epoch {checkpoint['epoch']} with loss {checkpoint['loss']:.4f}\")\n\n# Загрузка тестовых данных\ntest_df = load_test_data()\n\n# Создание тестового датасета\ntest_dataset = TestCellSegData(os.path.join(DATA_PATH, 'test'), \n                                test_df, test_augmentations)\ntest_loader = torch.utils.data.DataLoader(\n    test_dataset, batch_size=1, shuffle=False, num_workers=NUM_WORKERS)\n\n# Создание submission\nsubmission = create_submission(model, test_loader, device, threshold=0.5)\n\nprint(\"\\nFirst few rows of submission:\")\nprint(submission.head())\nprint(f\"\\nSubmission shape: {submission.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:46:08.122800Z","iopub.execute_input":"2026-03-30T22:46:08.123147Z","iopub.status.idle":"2026-03-30T22:46:10.134312Z","shell.execute_reply.started":"2026-03-30T22:46:08.123114Z","shell.execute_reply":"2026-03-30T22:46:10.133457Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Покажем несколько предсказаний на тестовых данных\nmodel.eval()\ntest_samples = list(test_dataset.ids[:3])\n\nfor img_id in test_samples:\n    # Загрузка изображения\n    img_path = os.path.join(DATA_PATH, 'test', f'{img_id}.png')\n    image = cv2.imread(img_path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    # Предсказание\n    img_tensor = torch.FloatTensor(image).permute(2, 0, 1).unsqueeze(0) / 255.0\n    img_tensor = img_tensor.to(device)\n    \n    with torch.no_grad():\n        predictions = model(img_tensor)[0]\n    \n    # Визуализация\n    plt.figure(figsize=(12, 8))\n    plt.imshow(image)\n    \n    masks = predictions['masks'].cpu().numpy()\n    scores = predictions['scores'].cpu().numpy()\n    \n    # Показываем только предсказания с уверенностью > 0.5\n    for mask, score in zip(masks, scores):\n        if score > 0.5:\n            binary_mask = (mask[0] > 0.5).astype(np.uint8)\n            # Создаем цветную маску\n            colored_mask = np.zeros((*binary_mask.shape, 3))\n            colored_mask[binary_mask > 0] = [0.3, 0.6, 0.9]\n            plt.imshow(colored_mask, alpha=0.4)\n    \n    plt.title(f\"Test Image: {img_id}\")\n    plt.axis('off')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:47:43.621127Z","iopub.execute_input":"2026-03-30T22:47:43.622101Z","iopub.status.idle":"2026-03-30T22:48:05.416604Z","shell.execute_reply.started":"2026-03-30T22:47:43.622054Z","shell.execute_reply":"2026-03-30T22:48:05.415795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Финальное создание submission файла\nsubmission = create_submission(model, test_loader, device, threshold=0.5)\n\n# Сохранение submission\nsubmission.to_csv('submission.csv', index=False)\nprint(\"\\nSubmission file saved as 'submission.csv'\")\n\n# Вывод статистики по предсказаниям\nprint(f\"\\nSubmission statistics:\")\nprint(f\"Total predictions: {len(submission)}\")\nprint(f\"Number of unique images: {submission['id'].nunique()}\")\nprint(f\"Cell type distribution:\\n{submission['cell_type'].value_counts()}\")\n\n# Отображение первых строк для проверки\nprint(\"\\nPreview of submission:\")\nprint(submission.head(10))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T22:48:57.344470Z","iopub.execute_input":"2026-03-30T22:48:57.344802Z","iopub.status.idle":"2026-03-30T22:48:58.943379Z","shell.execute_reply.started":"2026-03-30T22:48:57.344776Z","shell.execute_reply":"2026-03-30T22:48:58.942454Z"}},"outputs":[],"execution_count":null}]}