{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n!pip install -q -U kaggle\n\nos.environ[\"KAGGLE_USERNAME\"] = \"viollett\"\nos.environ[\"KAGGLE_KEY\"] = \"KGAT_c8752b6a2f62f0b68a6cc088ef8cd041\"\n\nprint(\"Kaggle credentials установлены\")\n!kaggle datasets list -m","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:46:54.504997Z","iopub.execute_input":"2026-09-02T10:46:54.505463Z","iopub.status.idle":"2026-09-02T10:47:01.269533Z","shell.execute_reply.started":"2026-09-02T10:46:54.505432Z","shell.execute_reply":"2026-09-02T10:47:01.268597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\n\nPREPROCESS_DIR = \"/kaggle/working/dasiam_preprocessed\"\nos.makedirs(os.path.join(PREPROCESS_DIR, \"got10k\"), exist_ok=True)\nos.makedirs(os.path.join(PREPROCESS_DIR, \"coco\"), exist_ok=True)\n\nsrc_dir = \"/kaggle/input/datasets/viollett/dasiamrpn-preprocessed\"  # исправленный путь\n\nfor root, _, files in os.walk(src_dir):\n    for fname in files:\n        src = os.path.join(root, fname)\n        rel = os.path.relpath(src, src_dir)\n        dst = os.path.join(PREPROCESS_DIR, rel)\n        os.makedirs(os.path.dirname(dst), exist_ok=True)\n        shutil.copy2(src, dst)\n\nprint(\"Восстановлено файлов:\", sum(len(f) for _, _, f in os.walk(PREPROCESS_DIR)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:48:00.757862Z","iopub.execute_input":"2026-09-02T10:48:00.758365Z","iopub.status.idle":"2026-09-02T10:48:05.400078Z","shell.execute_reply.started":"2026-09-02T10:48:00.758338Z","shell.execute_reply":"2026-09-02T10:48:05.399089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import re\n\ndef get_done_chunks(folder, prefix):\n    done = set()\n    if not os.path.exists(folder):\n        return done\n    for f in os.listdir(folder):\n        m = re.match(rf\"{prefix}_avg_part_(\\d+)\\.pkl\", f)\n        if m:\n            done.add(int(m.group(1)))\n    return done\n\ngot10k_done = get_done_chunks(os.path.join(PREPROCESS_DIR, \"got10k\"), \"got10k\")\ncoco_done = get_done_chunks(os.path.join(PREPROCESS_DIR, \"coco\"), \"coco\")\n\nprint(\"GOT10k готово:\", sorted(got10k_done))\nprint(\"COCO готово:\", sorted(coco_done))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:48:36.077502Z","iopub.execute_input":"2026-09-02T10:48:36.078369Z","iopub.status.idle":"2026-09-02T10:48:36.08508Z","shell.execute_reply.started":"2026-09-02T10:48:36.078338Z","shell.execute_reply":"2026-09-02T10:48:36.084253Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"total_videos = len([\n    f for f in os.listdir(os.path.join(GOT10K_PATH, \"train\"))\n    if f.startswith(\"GOT-10k_Train\")\n])\nprint(\"Всего видео в GOT10k train:\", total_videos)\nprint(\"Обработано чанками:\", len(got10k_done) * 500, \"(примерно, если chunk_size=500)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T11:39:13.122944Z","iopub.execute_input":"2026-08-30T11:39:13.124014Z","iopub.status.idle":"2026-08-30T11:39:13.270798Z","shell.execute_reply.started":"2026-08-30T11:39:13.123983Z","shell.execute_reply":"2026-08-30T11:39:13.269926Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"src_dir = \"/kaggle/input/datasets/viollett/dasiamrpn-preprocessed\"\n\nfor root, dirs, files in os.walk(src_dir):\n    for f in files:\n        print(os.path.join(root, f))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T11:39:18.822823Z","iopub.execute_input":"2026-08-30T11:39:18.82378Z","iopub.status.idle":"2026-08-30T11:39:18.831499Z","shell.execute_reply.started":"2026-08-30T11:39:18.823747Z","shell.execute_reply":"2026-08-30T11:39:18.830622Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pickle\nimport json\nimport cv2\nimport numpy as np\n\nfrom tqdm.auto import tqdm\n\n\nos.makedirs(\n    PREPROCESS_DIR,\n    exist_ok=True\n)\n\n\nKAGGLE_USERNAME = \"viollett\"\n\nKAGGLE_DATASET = (\n    f\"{KAGGLE_USERNAME}/dasiamrpn-preprocessed\"\n)\n\n\n\nGOT10K_CHUNK_SIZE = 500\n\nCOCO_CHUNK_SIZE = 5000\n\n\nmetadata_path = os.path.join(\n    PREPROCESS_DIR,\n    \"dataset-metadata.json\"\n)\n\nmetadata = {\n    \"title\": \"DaSiamRPN Preprocessed GOT10k COCO\",\n    \"id\": KAGGLE_DATASET,\n    \"licenses\": [\n        {\n            \"name\": \"CC0-1.0\"\n        }\n    ]\n}\n\nwith open(\n    metadata_path,\n    \"w\"\n) as f:\n\n    json.dump(\n        metadata,\n        f,\n        indent=4\n    )\n\n\nprint(\n    \"Preprocessing directory:\",\n    PREPROCESS_DIR\n)\n\nprint(\n    \"Kaggle Dataset:\",\n    KAGGLE_DATASET\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T11:39:34.743184Z","iopub.execute_input":"2026-08-30T11:39:34.74349Z","iopub.status.idle":"2026-08-30T11:39:34.753416Z","shell.execute_reply.started":"2026-08-30T11:39:34.743449Z","shell.execute_reply":"2026-08-30T11:39:34.752685Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\n\n\nos.makedirs(\n    os.path.join(PREPROCESS_DIR, \"got10k\"),\n    exist_ok=True\n)\n\nos.makedirs(\n    os.path.join(PREPROCESS_DIR, \"coco\"),\n    exist_ok=True\n)\n\nmetadata = {\n    \"title\": \"DaSiamRPN Preprocessed GOT10k COCO\",\n    \"id\": \"viollett/dasiamrpn-preprocessed\",\n    \"licenses\": [\n        {\n            \"name\": \"CC0-1.0\"\n        }\n    ]\n}\n\nwith open(\n    os.path.join(\n        PREPROCESS_DIR,\n        \"dataset-metadata.json\"\n    ),\n    \"w\"\n) as f:\n    json.dump(\n        metadata,\n        f,\n        indent=4\n    )\n\nprint(\"Structure created.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:02:06.862861Z","iopub.execute_input":"2026-09-02T10:02:06.863308Z","iopub.status.idle":"2026-09-02T10:02:06.871313Z","shell.execute_reply.started":"2026-09-02T10:02:06.863273Z","shell.execute_reply":"2026-09-02T10:02:06.870379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!kaggle datasets list -m","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T11:39:48.234607Z","iopub.execute_input":"2026-08-30T11:39:48.2356Z","iopub.status.idle":"2026-08-30T11:39:49.051384Z","shell.execute_reply.started":"2026-08-30T11:39:48.235566Z","shell.execute_reply":"2026-08-30T11:39:49.050356Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\n\n\ndef upload_to_kaggle_dataset(message):\n\n    print()\n    print(\"=\" * 70)\n    print(\"UPLOADING TO KAGGLE DATASET\")\n    print(\"=\" * 70)\n\n    result = subprocess.run(\n        [\n            \"kaggle\",\n            \"datasets\",\n            \"version\",\n            \"-p\", PREPROCESS_DIR,\n            \"-m\", message,\n            \"-r\", \"skip\",\n            \"--dir-mode\", \"zip\",   # <-- добавили\n        ],\n        capture_output=True,\n        text=True\n    )\n\n    print(result.stdout)\n\n    if result.returncode != 0:\n        print(result.stderr)\n        raise RuntimeError(\n            \"Kaggle Dataset upload failed.\"\n        )\n\n    print(\"Upload completed.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T11:39:56.033778Z","iopub.execute_input":"2026-08-30T11:39:56.034252Z","iopub.status.idle":"2026-08-30T11:39:56.040869Z","shell.execute_reply.started":"2026-08-30T11:39:56.0342Z","shell.execute_reply":"2026-08-30T11:39:56.03998Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"GOT10K_PATH = \"/kaggle/input/datasets/abhimanyukarshni/got10k\"\nCOCO_PATH = \"/kaggle/input/datasets/awsaf49/coco-2017-dataset/coco2017\"","metadata":{"_uuid":"75294057-c1f5-4cbe-a113-95e90ae2cef9","_cell_guid":"a127fd05-d090-4ae0-938d-f0e061963bfa","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-09-03T10:19:33.168634Z","iopub.execute_input":"2026-09-03T10:19:33.16956Z","iopub.status.idle":"2026-09-03T10:19:33.173158Z","shell.execute_reply.started":"2026-09-03T10:19:33.169514Z","shell.execute_reply":"2026-09-03T10:19:33.172266Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nprint(json.load(open(os.path.join(PREPROCESS_DIR, \"dataset-metadata.json\"))))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T11:40:00.207703Z","iopub.execute_input":"2026-08-30T11:40:00.208657Z","iopub.status.idle":"2026-08-30T11:40:00.214073Z","shell.execute_reply.started":"2026-08-30T11:40:00.208619Z","shell.execute_reply":"2026-08-30T11:40:00.213071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"result = subprocess.run(\n    [\n        \"kaggle\", \"datasets\", \"version\",\n        \"-p\", PREPROCESS_DIR,\n        \"-m\", \"test upload with dir-mode zip\",   # <-- любая строка вместо message\n        \"-r\", \"skip\",\n        \"--dir-mode\", \"zip\",\n    ],\n    capture_output=True, text=True\n)\nprint(result.stdout)\nprint(result.stderr)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T11:40:17.729255Z","iopub.execute_input":"2026-08-30T11:40:17.730301Z","iopub.status.idle":"2026-08-30T11:40:25.528878Z","shell.execute_reply.started":"2026-08-30T11:40:17.730265Z","shell.execute_reply":"2026-08-30T11:40:25.527974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader, ConcatDataset\n\nclass AnchorTargetGenerator:\n    \"\"\"\n    Класс для генерации анкоров и формирования целевых меток\n    классификации и регрессии Bounding Boxов.\n    \n    Класс преобразует Ground Truth координаты объекта на изображении в тензоры\n    целевых меток [Channels, Height, Width], готовых для подачи в функции потерь.\n    \"\"\"\n    def __init__(\n        self,\n        score_size: int = 19,\n        total_stride: int = 8,\n        search_size: int = 271,\n        anchor_ratios: list = [0.33, 0.5, 1.0, 2.0, 3.0],\n        anchor_scales: list = [8.0],\n        pos_iou_thresh: float = 0.6,\n        neg_iou_thresh: float = 0.3,\n    ):\n        self.score_size = score_size\n        self.total_stride = total_stride\n        self.search_size = search_size\n        self.anchor_ratios = anchor_ratios\n        self.anchor_scales = anchor_scales\n        self.pos_iou_thresh = pos_iou_thresh\n        self.neg_iou_thresh = neg_iou_thresh\n        \n        self.num_anchors_per_cell = len(anchor_ratios) * len(anchor_scales) # 5\n        \n        self.anchors = self._generate_all_anchors() # [1805, 4] -> (cx, cy, w, h)\n\n    def _generate_base_anchors(self):\n        \"\"\"\n        Создает набор базовых рамок-шаблонов (5 штук) для одной точки с центром в (0, 0).\n        Для каждого ratio вычисляются ширина w и высота h так, чтобы площадь рамки \n        оставалась постоянной (base_area), но менялась ее геометрия.\n        Возвращает: Массив базовых рамок формы [5, 4] в формате (x1, y1, x2, y2) \n        \"\"\"\n        base_anchors = []\n        base_area = (self.total_stride * self.anchor_scales[0]) ** 2\n        \n        for ratio in self.anchor_ratios:\n            w = np.sqrt(base_area / ratio)\n            h = w * ratio\n            base_anchors.append([-w / 2.0, -h / 2.0, w / 2.0, h / 2.0])\n            \n        return np.array(base_anchors, dtype=np.float32)\n\n    def _generate_all_anchors(self):\n        \"\"\"\n        Размножает 5 базовых шаблонов рамок по всей двумерной сетке 19x19, \n        привязывая их к абсолютным пиксельным координатам патча изображения 271x271.\n\n        Возвращает: Тензор всех анкоров кадра формы [1805, 4] в формате (cx, cy, w, h).\n        \"\"\"\n        base_anchors = self._generate_base_anchors() # [5, 4]\n        \n        ori = - (self.score_size // 2) * self.total_stride # -72\n        \n        grid_y = np.arange(self.score_size) * self.total_stride + ori\n        grid_x = np.arange(self.score_size) * self.total_stride + ori\n        \n        grid_y, grid_x = np.meshgrid(grid_y, grid_x, indexing='ij')\n        \n        center_offset = self.search_size / 2.0\n        grid_cx = grid_x + center_offset\n        grid_cy = grid_y + center_offset\n        \n        all_anchors = np.zeros((self.num_anchors_per_cell, self.score_size, self.score_size, 4), dtype=np.float32)\n        \n        for k, ba in enumerate(base_anchors):\n            w = ba[2] - ba[0]\n            h = ba[3] - ba[1]\n            all_anchors[k, :, :, 0] = grid_cx\n            all_anchors[k, :, :, 1] = grid_cy\n            all_anchors[k, :, :, 2] = w\n            all_anchors[k, :, :, 3] = h\n            \n        # [1805, 4]\n        return torch.tensor(all_anchors.reshape(-1, 4), dtype=torch.float32)\n\n    @staticmethod\n    def _compute_iou(anchors_xyxy, gt_xyxy):\n        \"\"\"\n        Статический метод для вычисления метрики Intersection over Union (IoU)\n        между всеми 1805 анкорами и одной целевой рамкой Ground Truth (GT).\n\n        Параметры:\n            anchors_xyxy: Анкоры кадра формы [1805, 4] в формате (x1, y1, x2, y2).\n            gt_xyxy: Рамка объекта формы [1, 4] в формате (x1, y1, x2, y2).\n\n        Возвращает: Вектор значений IoU формы [1805], содержащий числа от 0.0 до 1.0.\n        \"\"\"\n        inter_x1 = torch.max(anchors_xyxy[:, 0], gt_xyxy[0, 0])\n        inter_y1 = torch.max(anchors_xyxy[:, 1], gt_xyxy[0, 1])\n        inter_x2 = torch.min(anchors_xyxy[:, 2], gt_xyxy[0, 2])\n        inter_y2 = torch.min(anchors_xyxy[:, 3], gt_xyxy[0, 3])\n\n        inter_w = torch.clamp(inter_x2 - inter_x1, min=0)\n        inter_h = torch.clamp(inter_y2 - inter_y1, min=0)\n        inter_area = inter_w * inter_h\n\n        area_anchors = (anchors_xyxy[:, 2] - anchors_xyxy[:, 0]) * (anchors_xyxy[:, 3] - anchors_xyxy[:, 1])\n        area_gt = (gt_xyxy[0, 2] - gt_xyxy[0, 0]) * (gt_xyxy[0, 3] - gt_xyxy[0, 1])\n\n        union_area = area_anchors + area_gt - inter_area\n        return inter_area / (union_area + 1e-8)\n\n    def __call__(self, gt_bbox_patch: torch.Tensor, has_target: float):\n        \"\"\"\n        Основной метод вызова класса. Формирует целевые метки (таргеты) для обучения\n        нейросети на конкретном кадре/патче.\n\n        Параметры:\n            gt_bbox_patch (torch.Tensor): Тензор координат объекта (cx, cy, w, h).\n            has_target (float): Флаг присутствия объекта на кадре (1.0 — есть, 0.0 — нет).\n\n        Возвращает:\n                   cls_targets [10, 19, 19]: Метки классов (One-Hot: 2 канала * 5 анкоров).\n                   reg_targets [20, 19, 19]: Относительные сдвиги (4 координаты * 5 анкоров).\n                   reg_weights [20, 19, 19]: Весовые маски для расчета Loss регрессии.\n        \"\"\"\n        total_anchors = self.anchors.shape[0] # 1805\n        \n        cls_targets = torch.full((total_anchors,), -1, dtype=torch.long)\n        reg_targets = torch.zeros((total_anchors, 4), dtype=torch.float32)\n        reg_weights = torch.zeros((total_anchors, 4), dtype=torch.float32)\n\n        if has_target == 0.0:\n            cls_targets[:] = 0\n            keep_neg = torch.randperm(total_anchors)[:64]\n            ignore_mask = torch.ones(total_anchors, dtype=torch.bool)\n            ignore_mask[keep_neg] = False\n            cls_targets[ignore_mask] = -1\n        else:\n            anchors_cxcy = self.anchors\n            anchors_xyxy = torch.zeros_like(anchors_cxcy)\n            anchors_xyxy[:, 0] = anchors_cxcy[:, 0] - anchors_cxcy[:, 2] / 2.0\n            anchors_xyxy[:, 1] = anchors_cxcy[:, 1] - anchors_cxcy[:, 3] / 2.0\n            anchors_xyxy[:, 2] = anchors_cxcy[:, 0] + anchors_cxcy[:, 2] / 2.0\n            anchors_xyxy[:, 3] = anchors_cxcy[:, 1] + anchors_cxcy[:, 3] / 2.0\n\n            gt_cxcy = gt_bbox_patch.unsqueeze(0) # [1, 4]\n            gt_xyxy = torch.zeros_like(gt_cxcy)\n            gt_xyxy[0, 0] = gt_cxcy[0, 0] - gt_cxcy[0, 2] / 2.0\n            gt_xyxy[0, 1] = gt_cxcy[0, 1] - gt_cxcy[0, 3] / 2.0\n            gt_xyxy[0, 2] = gt_cxcy[0, 0] + gt_cxcy[0, 2] / 2.0\n            gt_xyxy[0, 3] = gt_cxcy[0, 1] + gt_cxcy[0, 3] / 2.0\n\n            ious = self._compute_iou(anchors_xyxy, gt_xyxy)\n\n            cls_targets[ious < self.neg_iou_thresh] = 0\n            cls_targets[ious >= self.pos_iou_thresh] = 1\n\n            max_iou_idx = torch.argmax(ious)\n            cls_targets[max_iou_idx] = 1\n            \n            pos_idx = torch.where(cls_targets == 1)[0]\n            neg_idx = torch.where(cls_targets == 0)[0]\n            \n            num_pos = min(len(pos_idx), 16)\n            num_neg = 64 - num_pos\n            \n            if len(pos_idx) > num_pos:\n                keep_pos = pos_idx[torch.randperm(len(pos_idx))[:num_pos]]\n                drop_pos = pos_idx[~torch.isin(pos_idx, keep_pos)]\n                cls_targets[drop_pos] = -1\n            \n            if len(neg_idx) > num_neg:\n                keep_neg = neg_idx[torch.randperm(len(neg_idx))[:num_neg]]\n                drop_neg = neg_idx[~torch.isin(neg_idx, keep_neg)]\n                cls_targets[drop_neg] = -1\n            pos_mask = (cls_targets == 1)\n            if pos_mask.sum() > 0:\n                pos_anchors = anchors_cxcy[pos_mask]\n                \n                dx = (gt_cxcy[0, 0] - pos_anchors[:, 0]) / pos_anchors[:, 2]\n                dy = (gt_cxcy[0, 1] - pos_anchors[:, 1]) / pos_anchors[:, 3]\n                dw = torch.log(torch.clamp(gt_cxcy[0, 2], min=1e-6) / torch.clamp(pos_anchors[:, 2], min=1e-6))\n                dh = torch.log(torch.clamp(gt_cxcy[0, 3], min=1e-6) / torch.clamp(pos_anchors[:, 3], min=1e-6))\n\n                reg_targets[pos_mask] = torch.stack([dx, dy, dw, dh], dim=1)\n                reg_weights[pos_mask] = 1.0\n\n        cls_targets_2ch = torch.zeros((total_anchors, 2), dtype=torch.float32)\n        cls_targets_2ch[cls_targets == 0, 0] = 1.0  # background\n        cls_targets_2ch[cls_targets == 1, 1] = 1.0  # foreground\n\n        cls_targets = cls_targets_2ch.view(self.num_anchors_per_cell, self.score_size, self.score_size, 2) \\\n                                     .permute(0, 3, 1, 2) \\\n                                     .reshape(self.num_anchors_per_cell * 2, self.score_size, self.score_size)\n\n        reg_targets = reg_targets.view(self.num_anchors_per_cell, self.score_size, self.score_size, 4) \\\n                                 .permute(0, 3, 1, 2) \\\n                                 .reshape(20, self.score_size, self.score_size)\n\n        reg_weights = reg_weights.view(self.num_anchors_per_cell, self.score_size, self.score_size, 4) \\\n                                 .permute(0, 3, 1, 2) \\\n                                 .reshape(20, self.score_size, self.score_size)\n\n        return cls_targets, reg_targets, reg_weights","metadata":{"_uuid":"521f7f8f-0394-4bca-86a4-eaff1577e53a","_cell_guid":"5f1026f8-5448-4f90-a803-f204befbee33","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-09-02T10:48:53.385706Z","iopub.execute_input":"2026-09-02T10:48:53.386153Z","iopub.status.idle":"2026-09-02T10:48:57.746361Z","shell.execute_reply.started":"2026-09-02T10:48:53.386126Z","shell.execute_reply":"2026-09-02T10:48:57.74563Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport random\nfrom abc import ABC, abstractmethod\n\nimport cv2\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\n\n\nclass BaseTrackingDataset(Dataset, ABC):\n\n    def __init__(\n        self,\n        template_size=127,\n        search_size=271,\n        context_amount=0.5,\n        swapRB=True,\n        augment=True,\n        template_shift_px=4,\n        search_shift_px=15,\n        scale_jitter_range=(0.9, 1.1),\n        flip_prob=0.5,\n        color_jitter_prob=0.3,\n        grayscale_prob=0.1,\n        blur_prob=0.05,\n        cutout_prob=0.1,\n    ):\n        \"\"\"\n        template_size:\n            Размер template patch, обычно 127x127.\n\n        search_size:\n            Размер search patch, обычно 271x271.\n\n        context_amount:\n            Количество контекста вокруг объекта.\n\n        swapRB:\n            BGR -> RGB.\n\n        augment:\n            Включить/выключить аугментации целиком.\n            Для validation/test рекомендуется False.\n\n        template_shift_px:\n            Максимальный сдвиг центра template в пикселях.\n\n        search_shift_px:\n            Максимальный сдвиг центра search region.\n\n        scale_jitter_range:\n            Случайное изменение размера search region.\n\n        flip_prob:\n            Вероятность horizontal flip.\n\n        color_jitter_prob:\n            Вероятность изменения brightness/contrast.\n\n        grayscale_prob:\n            Вероятность grayscale.\n\n        blur_prob:\n            Вероятность Gaussian blur.\n\n        cutout_prob:\n            Вероятность частичной окклюзии.\n        \"\"\"\n\n        super().__init__()\n\n\n        self.template_size = template_size\n        self.search_size = search_size\n        self.context_amount = context_amount\n        self.swapRB = swapRB\n\n\n        self.augment = augment\n\n        self.template_shift_px = template_shift_px\n        self.search_shift_px = search_shift_px\n\n        self.scale_jitter_range = scale_jitter_range\n\n        self.flip_prob = flip_prob\n\n        self.color_jitter_prob = color_jitter_prob\n        self.grayscale_prob = grayscale_prob\n        self.blur_prob = blur_prob\n        self.cutout_prob = cutout_prob\n\n\n        self.anchor_target_generator = AnchorTargetGenerator(\n            score_size=19,\n            total_stride=8,\n            search_size=search_size,\n            pos_iou_thresh=0.6,\n            neg_iou_thresh=0.3,\n        )\n    def getSubwindow(\n        self,\n        img,\n        targetBox,\n        originalSize,\n        avgChans,\n    ):\n        \"\"\"\n        Вырезает квадратный crop вокруг центра targetBox.\n\n        targetBox:\n            (center_x, center_y, width, height)\n\n        Если crop выходит за границы изображения,\n        используется padding средним цветом изображения.\n        \"\"\"\n        avgChans = tuple(float(c) for c in avgChans) \n        img_h, img_w = img.shape[:2]\n\n        center_x, center_y, width, height = targetBox\n\n        c = (originalSize + 1) / 2\n\n        xMin = int(np.round(center_x - c))\n        yMin = int(np.round(center_y - c))\n\n        xMax = xMin + int(originalSize)\n        yMax = yMin + int(originalSize)\n\n        leftPad = max(0, -xMin)\n        topPad = max(0, -yMin)\n\n        rightPad = max(0, xMax - img_w)\n        bottomPad = max(0, yMax - img_h)\n\n        xMin += leftPad\n        xMax += leftPad\n\n        yMin += topPad\n        yMax += topPad\n\n\n        if (\n            leftPad == 0\n            and topPad == 0\n            and rightPad == 0\n            and bottomPad == 0\n        ):\n            crop = img[\n                yMin:yMax,\n                xMin:xMax\n            ].copy()\n\n\n        else:\n\n            padded = cv2.copyMakeBorder(\n                img,\n                topPad,\n                bottomPad,\n                leftPad,\n                rightPad,\n                cv2.BORDER_CONSTANT,\n                value=avgChans,\n            )\n\n            crop = padded[\n                yMin:yMax,\n                xMin:xMax\n            ].copy()\n\n        return crop\n\n\n    def _preprocess_image(\n        self,\n        image,\n        target_size,\n        avgChans=None,\n    ):\n        \"\"\"\n        Resize\n        BGR -> RGB\n        subtract mean\n        HWC -> CHW\n        \"\"\"\n\n        h, w = target_size\n\n\n        if image.shape[:2] != (h, w):\n\n            image = cv2.resize(\n                image,\n                (w, h),\n                interpolation=cv2.INTER_LINEAR,\n            )\n\n\n        if self.swapRB:\n\n            image = cv2.cvtColor(\n                image,\n                cv2.COLOR_BGR2RGB,\n            )\n\n\n        image = image.astype(np.float32)\n\n\n        if avgChans is not None:\n\n            image = image - np.asarray(\n                avgChans,\n                dtype=np.float32,\n            )\n\n\n        image = torch.from_numpy(\n            image\n        ).permute(2, 0, 1).float()\n\n        return image\n\n\n    @staticmethod\n    def _color_jitter(\n        image,\n        prob=0.3,\n    ):\n  \n\n        if random.random() >= prob:\n            return image\n\n        alpha = random.uniform(\n            0.8,\n            1.2,\n        )\n\n        beta = random.uniform(\n            -15,\n            15,\n        )\n\n        image = cv2.convertScaleAbs(\n            image,\n            alpha=alpha,\n            beta=beta,\n        )\n\n        return image\n\n    @staticmethod\n    def _grayscale(\n        image,\n        prob=0.1,\n    ):\n        \"\"\"\n        Перевод изображения в grayscale,\n        сохраняя 3 канала.\n        \"\"\"\n\n        if random.random() >= prob:\n            return image\n\n        gray = cv2.cvtColor(\n            image,\n            cv2.COLOR_BGR2GRAY,\n        )\n\n        image = cv2.cvtColor(\n            gray,\n            cv2.COLOR_GRAY2BGR,\n        )\n\n        return image\n\n    @staticmethod\n    def _blur(\n        image,\n        prob=0.05,\n    ):\n        \"\"\"\n        Небольшой Gaussian blur.\n        \"\"\"\n\n        if random.random() >= prob:\n            return image\n\n        kernel = random.choice(\n            [3, 5]\n        )\n\n        image = cv2.GaussianBlur(\n            image,\n            (kernel, kernel),\n            0,\n        )\n\n        return image\n\n    @staticmethod\n    def _cutout(\n        image,\n        avg,\n        prob=0.1,\n    ):\n        \"\"\"\n        Частичная окклюзия объекта.\n\n        Вырезается небольшой прямоугольник\n        и заполняется средним цветом.\n        \"\"\"\n\n        if random.random() >= prob:\n            return image\n\n        h, w = image.shape[:2]\n\n        # Размер окклюзии:\n        # примерно 10-20% изображения\n        cut_w = random.randint(\n            max(1, w // 10),\n            max(2, w // 5),\n        )\n\n        cut_h = random.randint(\n            max(1, h // 10),\n            max(2, h // 5),\n        )\n\n        x = random.randint(\n            0,\n            max(0, w - cut_w),\n        )\n\n        y = random.randint(\n            0,\n            max(0, h - cut_h),\n        )\n\n        image = image.copy()\n\n        image[\n            y:y + cut_h,\n            x:x + cut_w\n        ] = np.asarray(\n            avg,\n            dtype=np.uint8,\n        )\n\n        return image\n\n    def _apply_photometric_augmentation(\n        self,\n        image,\n        avg,\n    ):\n        \"\"\"\n        Все photometric augmentations.\n\n        Они НЕ меняют bbox.\n        \"\"\"\n\n        image = self._color_jitter(\n            image,\n            self.color_jitter_prob,\n        )\n\n        image = self._grayscale(\n            image,\n            self.grayscale_prob,\n        )\n\n        image = self._blur(\n            image,\n            self.blur_prob,\n        )\n\n        image = self._cutout(\n            image,\n            avg,\n            self.cutout_prob,\n        )\n\n        return image\n\n\n    def _maybe_flip(\n        self,\n        crop,\n        gt_bbox_patch,\n    ):\n        \"\"\"\n        Horizontal flip search crop.\n\n        gt_bbox_patch:\n            [cx, cy, w, h]\n        \"\"\"\n\n        if random.random() >= self.flip_prob:\n\n            return crop, gt_bbox_patch\n\n\n        crop = cv2.flip(\n            crop,\n            1,\n        )\n\n\n        gt_bbox_patch = gt_bbox_patch.copy()\n\n        gt_bbox_patch[0] = (\n            self.search_size\n            - 1\n            - gt_bbox_patch[0]\n        )\n\n        return crop, gt_bbox_patch\n\n\n    def _crop(self, image, bbox, output_size, avg=None):\n        x, y, w, h = bbox\n        cx = x + w / 2\n        cy = y + h / 2\n\n        if self.augment and self.template_shift_px > 0:\n            cx += np.random.uniform(-self.template_shift_px, self.template_shift_px)\n            cy += np.random.uniform(-self.template_shift_px, self.template_shift_px)\n\n        if avg is None:\n            avg = np.mean(image, axis=(0, 1)).tolist()   # fallback, если avg не передали\n\n        wc = w + self.context_amount * (w + h)\n        hc = h + self.context_amount * (w + h)\n        sz = np.sqrt(wc * hc)\n\n        crop = self.getSubwindow(image, (cx, cy, w, h), sz, avg)\n\n        if self.augment:\n            crop = self._apply_photometric_augmentation(crop, avg)\n\n        crop_avg = np.mean(crop, axis=(0, 1)).tolist()  # это mean по маленькому патчу — дёшево, не трогаем\n        tensor = self._preprocess_image(crop, (output_size, output_size), crop_avg)\n\n        return tensor, sz, crop_avg\n\n\n    def _build_search(self, image, bbox, avg=None):\n        x, y, w, h = bbox\n\n        if avg is None:\n            avg = np.mean(image, axis=(0, 1)).tolist()   # fallback\n\n        wc = w + self.context_amount * (w + h)\n        hc = h + self.context_amount * (w + h)\n        sz = np.sqrt(wc * hc)\n\n        scale_z = self.template_size / sz\n        search_pad = (self.search_size - self.template_size) / 2\n        pad = search_pad / scale_z\n        sx = sz + 2 * pad\n\n        if self.augment:\n            scale_min, scale_max = self.scale_jitter_range\n            sx *= np.random.uniform(scale_min, scale_max)\n\n        shift_x = shift_y = 0.0\n        if self.augment:\n            shift_x = np.random.uniform(-self.search_shift_px, self.search_shift_px)\n            shift_y = np.random.uniform(-self.search_shift_px, self.search_shift_px)\n\n        center_x = x + w / 2 + shift_x\n        center_y = y + h / 2 + shift_y\n\n        crop = self.getSubwindow(image, (center_x, center_y, w, h), sx, avg)\n\n        if self.augment:\n            crop = self._apply_photometric_augmentation(crop, avg)\n\n        scale = self.search_size / sx\n        gt_cx = (x + w / 2 - center_x) * scale + self.search_size / 2\n        gt_cy = (y + h / 2 - center_y) * scale + self.search_size / 2\n        gt_w = max(1.0, w * scale)\n        gt_h = max(1.0, h * scale)\n\n        gt_bbox_patch = np.array([gt_cx, gt_cy, gt_w, gt_h], dtype=np.float32)\n\n        if self.augment:\n            crop, gt_bbox_patch = self._maybe_flip(crop, gt_bbox_patch)\n\n        crop_avg = np.mean(crop, axis=(0, 1)).tolist()\n        tensor = self._preprocess_image(crop, (self.search_size, self.search_size), crop_avg)\n\n        return tensor, gt_bbox_patch\n\n\n    def build_sample(self, template_img, template_bbox, search_img, search_bbox,\n                      has_target, template_avg=None, search_avg=None):\n\n        template, _, _ = self._crop(template_img, template_bbox, self.template_size, avg=template_avg)\n        search, gt_bbox_patch = self._build_search(search_img, search_bbox, avg=search_avg)\n\n        gt_tensor = torch.from_numpy(gt_bbox_patch)\n        cls_targets, reg_targets, reg_weights = self.anchor_target_generator(gt_tensor, has_target)\n\n        return {\n            \"template\": template,\n            \"search\": search,\n            \"gt_bbox_patch\": gt_tensor,\n            \"has_target\": torch.tensor(has_target, dtype=torch.float32),\n            \"cls_targets\": cls_targets,\n            \"reg_targets\": reg_targets,\n            \"reg_weights\": reg_weights,\n        }\n\n\n\n    @abstractmethod\n    def __getitem__(self, index):\n        pass\n","metadata":{"_uuid":"b79b09ea-ccfe-41ae-843d-d89aa189d518","_cell_guid":"4688ee3e-3fde-46ab-a0ff-1d64c101c534","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-09-02T10:48:58.494686Z","iopub.execute_input":"2026-09-02T10:48:58.495356Z","iopub.status.idle":"2026-09-02T10:48:58.73175Z","shell.execute_reply.started":"2026-09-02T10:48:58.495324Z","shell.execute_reply":"2026-09-02T10:48:58.730883Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_got10k_chunks(\n    got10k_dir,\n    output_dir=PREPROCESS_DIR,\n    chunk_size=500,\n):\n    \n    os.makedirs(\n        output_dir,\n        exist_ok=True\n    )\n\n    train_dir = os.path.join(\n        got10k_dir,\n        \"train\"\n    )\n\n\n    video_folders = sorted(\n        os.path.join(\n            train_dir,\n            f\n        )\n        for f in os.listdir(train_dir)\n        if f.startswith(\"GOT-10k_Train\")\n        and os.path.isdir(\n            os.path.join(\n                train_dir,\n                f\n            )\n        )\n    )\n\n    print(\n        f\"[GOT10k] Videos: {len(video_folders)}\"\n    )\n\n\n    metadata_path = os.path.join(\n        output_dir,\n        \"got10k_metadata.pkl\"\n    )\n\n    if os.path.exists(metadata_path):\n\n        print(\n            \"[GOT10k] Metadata already exists.\"\n        )\n\n        with open(\n            metadata_path,\n            \"rb\"\n        ) as f:\n\n            metadata = pickle.load(f)\n\n    else:\n\n        videos = []\n\n        for video_dir in tqdm(\n            video_folders,\n            desc=\"Reading GOT10k metadata\"\n        ):\n\n            gt_path = os.path.join(\n                video_dir,\n                \"groundtruth.txt\"\n            )\n\n            bboxes = []\n\n            with open(\n                gt_path,\n                \"r\"\n            ) as f:\n\n                for line in f:\n\n                    line = line.strip()\n\n                    if not line:\n                        continue\n\n                    line = line.replace(\n                        \"\\t\",\n                        \",\"\n                    )\n\n                    bbox = [\n                        float(x)\n                        for x in line.split(\",\")\n                    ]\n\n                    bboxes.append(\n                        bbox\n                    )\n\n            image_paths = [\n                os.path.join(\n                    video_dir,\n                    f\"{i + 1:08d}.jpg\"\n                )\n                for i in range(\n                    len(bboxes)\n                )\n            ]\n\n            videos.append(\n                {\n                    \"video_dir\":\n                        video_dir,\n\n                    \"image_paths\":\n                        image_paths,\n\n                    \"bboxes\":\n                        np.asarray(\n                            bboxes,\n                            dtype=np.float32\n                        )\n                }\n            )\n\n        metadata = {\n            \"videos\": videos\n        }\n\n        with open(\n            metadata_path,\n            \"wb\"\n        ) as f:\n\n            pickle.dump(\n                metadata,\n                f,\n                protocol=pickle.HIGHEST_PROTOCOL\n            )\n\n        print(\n            f\"[SAVED] {metadata_path}\"\n        )\n\n    videos = metadata[\"videos\"]\n\n\n    total_videos = len(videos)\n\n    num_chunks = (\n        total_videos +\n        chunk_size -\n        1\n    ) // chunk_size\n\n    print(\n        f\"[GOT10k] Chunks: {num_chunks}\"\n    )\n\n    for chunk_id in range(\n        num_chunks\n    ):\n\n        chunk_path = os.path.join(\n            output_dir,\n            f\"got10k_avg_part_{chunk_id:03d}.pkl\"\n        )\n\n\n        if os.path.exists(\n            chunk_path\n        ):\n\n            print(\n                f\"[SKIP] GOT10k chunk \"\n                f\"{chunk_id + 1}/{num_chunks}\"\n            )\n\n            continue\n\n        start = (\n            chunk_id *\n            chunk_size\n        )\n\n        end = min(\n            start + chunk_size,\n            total_videos\n        )\n\n        print()\n        print(\"=\" * 70)\n        print(\n            f\"GOT10k chunk \"\n            f\"{chunk_id + 1}/{num_chunks}\"\n        )\n        print(\n            f\"Videos: {start} -> {end - 1}\"\n        )\n        print(\"=\" * 70)\n\n        image_avgs = {}\n\n\n        for video in tqdm(\n            videos[start:end],\n            desc=f\"GOT chunk {chunk_id}\"\n        ):\n\n            for image_path in video[\n                \"image_paths\"\n            ]:\n\n                if not os.path.exists(\n                    image_path\n                ):\n\n                    continue\n\n                img = cv2.imread(\n                    image_path\n                )\n\n                if img is None:\n\n                    print(\n                        f\"[WARNING] Cannot read: \"\n                        f\"{image_path}\"\n                    )\n\n                    continue\n\n                avg = np.mean(\n                    img,\n                    axis=(0, 1),\n                    dtype=np.float32\n                )\n\n                image_avgs[\n                    image_path\n                ] = avg\n\n\n        with open(\n            chunk_path,\n            \"wb\"\n        ) as f:\n\n            pickle.dump(\n                image_avgs,\n                f,\n                protocol=pickle.HIGHEST_PROTOCOL\n            )\n\n        print(\n            f\"[SAVED] {chunk_path}\"\n        )\n\n\n        upload_to_kaggle_dataset(\n            message=(\n                f\"GOT10k avg chunk \"\n                f\"{chunk_id:03d}\"\n            )\n        )\n\n    print()\n    print(\n        \"[GOT10k] preprocessing finished.\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:11:26.83151Z","iopub.execute_input":"2026-09-02T10:11:26.831884Z","iopub.status.idle":"2026-09-02T10:11:26.851364Z","shell.execute_reply.started":"2026-09-02T10:11:26.831851Z","shell.execute_reply":"2026-09-02T10:11:26.850228Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pycocotools.coco import COCO\n\n\ndef preprocess_coco_chunks(\n    coco_dir,\n    output_dir=PREPROCESS_DIR,\n    chunk_size=5000,\n):\n\n    os.makedirs(\n        output_dir,\n        exist_ok=True\n    )\n\n    image_dir = os.path.join(\n        coco_dir,\n        \"train2017\"\n    )\n\n    ann_file = os.path.join(\n        coco_dir,\n        \"annotations\",\n        \"instances_train2017.json\"\n    )\n\n\n    coco = COCO(\n        ann_file\n    )\n\n\n    metadata_path = os.path.join(\n        output_dir,\n        \"coco_metadata.pkl\"\n    )\n\n    if os.path.exists(\n        metadata_path\n    ):\n\n        print(\n            \"[COCO] Metadata already exists.\"\n        )\n\n        with open(\n            metadata_path,\n            \"rb\"\n        ) as f:\n\n            metadata = pickle.load(f)\n\n    else:\n\n        image_ids = []\n        objects = []\n        objects_by_category = {}\n\n        print(\n            \"Reading COCO annotations...\"\n        )\n\n        for img_id in tqdm(\n            coco.getImgIds(),\n            desc=\"COCO annotations\"\n        ):\n\n            anns = coco.loadAnns(\n                coco.getAnnIds(\n                    imgIds=[img_id],\n                    iscrowd=False\n                )\n            )\n\n            valid_objects = []\n\n            for ann in anns:\n\n                x, y, w, h = ann[\n                    \"bbox\"\n                ]\n\n                if w <= 0 or h <= 0:\n                    continue\n\n                obj_idx = len(\n                    objects\n                )\n\n                obj = {\n                    \"image_id\":\n                        img_id,\n\n                    \"bbox\":\n                        [\n                            x,\n                            y,\n                            w,\n                            h\n                        ],\n\n                    \"category_id\":\n                        ann[\"category_id\"],\n\n                    \"area\":\n                        ann[\"area\"],\n\n                    \"ratio\":\n                        w / h,\n                }\n\n                objects.append(\n                    obj\n                )\n\n                valid_objects.append(\n                    obj_idx\n                )\n\n                objects_by_category.setdefault(\n                    ann[\"category_id\"],\n                    []\n                ).append(\n                    obj_idx\n                )\n\n            if valid_objects:\n\n                image_ids.append(\n                    img_id\n                )\n\n\n        objects_by_image = {}\n\n        for obj_idx, obj in enumerate(\n            objects\n        ):\n\n            image_id = obj[\n                \"image_id\"\n            ]\n\n            objects_by_image.setdefault(\n                image_id,\n                []\n            ).append(\n                obj_idx\n            )\n\n\n        image_paths = {}\n\n        for image_id in tqdm(\n            image_ids,\n            desc=\"COCO image paths\"\n        ):\n\n            info = coco.loadImgs(\n                image_id\n            )[0]\n\n            image_paths[\n                image_id\n            ] = os.path.join(\n                image_dir,\n                info[\"file_name\"]\n            )\n\n        metadata = {\n\n            \"image_ids\":\n                image_ids,\n\n            \"image_paths\":\n                image_paths,\n\n            \"objects\":\n                objects,\n\n            \"objects_by_image\":\n                objects_by_image,\n\n            \"objects_by_category\":\n                objects_by_category,\n        }\n\n        with open(\n            metadata_path,\n            \"wb\"\n        ) as f:\n\n            pickle.dump(\n                metadata,\n                f,\n                protocol=pickle.HIGHEST_PROTOCOL\n            )\n\n        print(\n            f\"[SAVED] {metadata_path}\"\n        )\n\n    image_ids = metadata[\n        \"image_ids\"\n    ]\n\n    image_paths = metadata[\n        \"image_paths\"\n    ]\n\n\n    total_images = len(\n        image_ids\n    )\n\n    num_chunks = (\n        total_images +\n        chunk_size -\n        1\n    ) // chunk_size\n\n    print(\n        f\"[COCO] Images: {total_images}\"\n    )\n\n    print(\n        f\"[COCO] Chunks: {num_chunks}\"\n    )\n\n\n    for chunk_id in range(\n        num_chunks\n    ):\n\n        chunk_path = os.path.join(\n            output_dir,\n            f\"coco_avg_part_{chunk_id:03d}.pkl\"\n        )\n\n\n        if os.path.exists(\n            chunk_path\n        ):\n\n            print(\n                f\"[SKIP] COCO chunk \"\n                f\"{chunk_id + 1}/{num_chunks}\"\n            )\n\n            continue\n\n        start = (\n            chunk_id *\n            chunk_size\n        )\n\n        end = min(\n            start + chunk_size,\n            total_images\n        )\n\n        chunk_ids = image_ids[\n            start:end\n        ]\n\n        print()\n        print(\"=\" * 70)\n        print(\n            f\"COCO chunk \"\n            f\"{chunk_id + 1}/{num_chunks}\"\n        )\n        print(\n            f\"Images: {start} -> {end - 1}\"\n        )\n        print(\n            f\"Count: {len(chunk_ids)}\"\n        )\n        print(\"=\" * 70)\n\n        image_avgs = {}\n\n\n        for image_id in tqdm(\n            chunk_ids,\n            desc=f\"COCO chunk {chunk_id}\"\n        ):\n\n            path = image_paths[\n                image_id\n            ]\n\n            img = cv2.imread(\n                path\n            )\n\n            if img is None:\n\n                print(\n                    f\"[WARNING] Cannot read: \"\n                    f\"{path}\"\n                )\n\n                continue\n\n            avg = np.mean(\n                img,\n                axis=(0, 1),\n                dtype=np.float32\n            )\n\n            image_avgs[\n                image_id\n            ] = avg\n\n\n        with open(\n            chunk_path,\n            \"wb\"\n        ) as f:\n\n            pickle.dump(\n                image_avgs,\n                f,\n                protocol=pickle.HIGHEST_PROTOCOL\n            )\n\n        print(\n            f\"[SAVED] {chunk_path}\"\n        )\n\n\n        upload_to_kaggle_dataset(\n            message=(\n                f\"COCO avg chunk \"\n                f\"{chunk_id:03d}\"\n            )\n        )\n\n    print()\n    print(\n        \"[COCO] preprocessing finished.\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:11:33.40913Z","iopub.execute_input":"2026-09-02T10:11:33.410467Z","iopub.status.idle":"2026-09-02T10:11:33.450721Z","shell.execute_reply.started":"2026-09-02T10:11:33.410418Z","shell.execute_reply":"2026-09-02T10:11:33.449595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preprocess_got10k_chunks(\n    got10k_dir=GOT10K_PATH,\n    output_dir=os.path.join(PREPROCESS_DIR, \"got10k\"),\n    chunk_size=500\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-29T16:31:06.441665Z","iopub.execute_input":"2026-08-29T16:31:06.44206Z","iopub.status.idle":"2026-08-29T19:31:50.124012Z","shell.execute_reply.started":"2026-08-29T16:31:06.44203Z","shell.execute_reply":"2026-08-29T19:31:50.122865Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preprocess_coco_chunks(\n    coco_dir=COCO_PATH,\n    output_dir=os.path.join(PREPROCESS_DIR, \"coco\"),\n    chunk_size=5000\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T11:53:29.542693Z","iopub.execute_input":"2026-08-30T11:53:29.543874Z","iopub.status.idle":"2026-08-30T12:58:45.048492Z","shell.execute_reply.started":"2026-08-30T11:53:29.543838Z","shell.execute_reply":"2026-08-30T12:58:45.046689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!kaggle datasets create \\\n    -p /kaggle/working/dasiam_preprocessed \\\n    -r zip","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pickle\nimport random\n\nimport cv2\nimport numpy as np\n\n\nclass GOT10kDataset(BaseTrackingDataset):\n\n    def __init__(\n        self,\n        got10k_dir,\n        metadata_path=\"/kaggle/working/dasiam_preprocessed/got10k/got10k_metadata.pkl\",\n        avg_dir=\"/kaggle/working/dasiam_preprocessed/got10k\",\n        avg_chunk_size=500,\n        template_size=127,\n        search_size=271,\n        neg_pair_prob=0.2,\n        swapRB=True,\n    ):\n\n        super().__init__(\n            template_size=template_size,\n            search_size=search_size,\n            context_amount=0.5,\n            swapRB=swapRB,\n        )\n\n        self.got10k_dir = got10k_dir\n        self.neg_pair_prob = neg_pair_prob\n\n        self.avg_dir = avg_dir\n        self.avg_chunk_size = avg_chunk_size\n\n\n        with open(\n            metadata_path,\n            \"rb\"\n        ) as f:\n\n            data = pickle.load(f)\n\n        self.videos = data[\"videos\"]\n\n        print(\n            f\"[GOT10k] Videos: \"\n            f\"{len(self.videos)}\"\n        )\n\n\n        self._avg_cache = {}\n\n        # путь изображения -> индекс видео\n        self._image_to_chunk = {}\n\n        for video_idx, video in enumerate(\n            self.videos\n        ):\n\n            chunk_id = (\n                video_idx\n                // self.avg_chunk_size\n            )\n\n            for image_path in video[\n                \"image_paths\"\n            ]:\n\n                self._image_to_chunk[\n                    image_path\n                ] = chunk_id\n\n    def __len__(self):\n\n        return len(\n            self.videos\n        )\n\n\n    def _get_avg(self, image_path):\n\n        chunk_id = self._image_to_chunk[\n            image_path\n        ]\n\n        # chunk уже загружен\n        if chunk_id not in self._avg_cache:\n\n            chunk_path = os.path.join(\n                self.avg_dir,\n                f\"got10k_avg_part_{chunk_id:03d}.pkl\"\n            )\n\n            with open(\n                chunk_path,\n                \"rb\"\n            ) as f:\n\n                self._avg_cache[\n                    chunk_id\n                ] = pickle.load(f)\n\n        return self._avg_cache[\n            chunk_id\n        ][image_path]\n\n\n    @staticmethod\n    def _load_image(path):\n\n        img = cv2.imread(\n            path\n        )\n\n        if img is None:\n\n            raise RuntimeError(\n                f\"Не удалось загрузить: \"\n                f\"{path}\"\n            )\n\n        return img\n\n\n    def __getitem__(self, index):\n\n\n        video_z = self.videos[\n            index\n        ]\n\n        image_paths_z = video_z[\n            \"image_paths\"\n        ]\n\n        bboxes_z = video_z[\n            \"bboxes\"\n        ]\n\n        num_frames = len(\n            image_paths_z\n        )\n\n        idx_z = random.randrange(\n            num_frames\n        )\n\n        template_path = (\n            image_paths_z[idx_z]\n        )\n\n        template_bbox = (\n            bboxes_z[idx_z].tolist()\n        )\n\n        template_img = self._load_image(\n            template_path\n        )\n\n        template_avg = self._get_avg(\n            template_path\n        )\n\n\n        negative = (\n            random.random()\n            < self.neg_pair_prob\n        )\n\n        if negative:\n\n\n            neg_video_idx = random.randrange(\n                len(self.videos)\n            )\n\n            while neg_video_idx == index:\n\n                neg_video_idx = random.randrange(\n                    len(self.videos)\n                )\n\n            video_x = self.videos[\n                neg_video_idx\n            ]\n\n            idx_x = random.randrange(\n                len(video_x[\"image_paths\"])\n            )\n\n            search_path = (\n                video_x[\"image_paths\"][idx_x]\n            )\n\n            search_bbox = (\n                video_x[\"bboxes\"][idx_x].tolist()\n            )\n\n            has_target = 0.0\n\n        else:\n\n            video_x = video_z\n\n            min_idx = max(\n                0,\n                idx_z - 100\n            )\n\n            max_idx = min(\n                num_frames - 1,\n                idx_z + 100\n            )\n\n            idx_x = random.randint(\n                min_idx,\n                max_idx\n            )\n\n            search_path = (\n                image_paths_z[idx_x]\n            )\n\n            search_bbox = (\n                bboxes_z[idx_x].tolist()\n            )\n\n            has_target = 1.0\n\n\n        search_img = self._load_image(\n            search_path\n        )\n\n        search_avg = self._get_avg(\n            search_path\n        )\n\n\n        sample = self.build_sample(\n\n            template_img=template_img,\n\n            template_bbox=template_bbox,\n\n            search_img=search_img,\n\n            search_bbox=search_bbox,\n\n            has_target=has_target,\n\n            template_avg=template_avg,\n\n            search_avg=search_avg,\n        )\n\n        sample[\"dataset\"] = \"GOT10k\"\n\n        sample[\"pair_type\"] = (\n            \"negative\"\n            if has_target == 0.0\n            else \"positive\"\n        )\n\n        return sample","metadata":{"_uuid":"fb4cd871-15bb-4aae-a0c9-a241200c648b","_cell_guid":"cc16f86b-4c6f-4a3f-be2d-74ed1aafc6f5","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-09-02T10:49:37.210793Z","iopub.execute_input":"2026-09-02T10:49:37.211062Z","iopub.status.idle":"2026-09-02T10:49:37.226529Z","shell.execute_reply.started":"2026-09-02T10:49:37.211041Z","shell.execute_reply":"2026-09-02T10:49:37.225549Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pickle\nimport random\n\nimport cv2\nimport numpy as np\nfrom functools import lru_cache\n\n\nclass COCODataset(BaseTrackingDataset):\n\n    def __init__(\n        self,\n        coco_dir,\n        metadata_path=\"/kaggle/working/dasiam_preprocessed/coco/coco_metadata.pkl\",\n        avg_dir=\"/kaggle/working/dasiam_preprocessed/coco\",\n        avg_chunk_size=5000,\n        template_size=127,\n        search_size=271,\n        neg_pair_prob=0.2,\n        hard_negative_prob=0.7,\n        swapRB=True,\n        image_cache_size=256,\n    ):\n\n        super().__init__(\n            template_size=template_size,\n            search_size=search_size,\n            context_amount=0.5,\n            swapRB=swapRB,\n        )\n\n        self.coco_dir = coco_dir\n\n        self.neg_pair_prob = (\n            neg_pair_prob\n        )\n\n        self.hard_negative_prob = (\n            hard_negative_prob\n        )\n\n        self.avg_dir = avg_dir\n        self.avg_chunk_size = (\n            avg_chunk_size\n        )\n\n\n        print(\n            \"[COCO] Loading metadata...\"\n        )\n\n        with open(\n            metadata_path,\n            \"rb\"\n        ) as f:\n\n            data = pickle.load(f)\n\n        self.image_ids = data[\n            \"image_ids\"\n        ]\n\n        self.image_paths = data[\n            \"image_paths\"\n        ]\n\n        self.objects = data[\n            \"objects\"\n        ]\n\n        self.objects_by_image = data[\n            \"objects_by_image\"\n        ]\n\n        self.objects_by_category = data[\n            \"objects_by_category\"\n        ]\n\n        print(\n            f\"[COCO] Images: \"\n            f\"{len(self.image_ids)}\"\n        )\n\n        print(\n            f\"[COCO] Objects: \"\n            f\"{len(self.objects)}\"\n        )\n\n\n        self._image_index = {\n            image_id: i\n            for i, image_id in enumerate(\n                self.image_ids\n            )\n        }\n\n\n        self._avg_cache = {}\n\n\n        self._cat_arrays = {}\n\n        for cat_id, idxs in (\n            self.objects_by_category.items()\n        ):\n\n            ratios = np.array(\n                [\n                    self.objects[i][\"ratio\"]\n                    for i in idxs\n                ],\n                dtype=np.float32\n            )\n\n            areas = np.array(\n                [\n                    self.objects[i][\"area\"]\n                    for i in idxs\n                ],\n                dtype=np.float32\n            )\n\n            image_ids_arr = np.array(\n                [\n                    self.objects[i][\"image_id\"]\n                    for i in idxs\n                ],\n                dtype=np.int64\n            )\n\n            self._cat_arrays[\n                cat_id\n            ] = {\n\n                \"idxs\":\n                    np.asarray(\n                        idxs,\n                        dtype=np.int32\n                    ),\n\n                \"ratios\":\n                    ratios,\n\n                \"areas\":\n                    areas,\n\n                \"image_ids\":\n                    image_ids_arr,\n            }\n\n\n        self._load_image_cached = (\n            lru_cache(\n                maxsize=image_cache_size\n            )(\n                self._load_image_from_disk\n            )\n        )\n\n\n    def _get_avg(self, image_id):\n\n        index = self._image_index[\n            image_id\n        ]\n\n        chunk_id = (\n            index\n            // self.avg_chunk_size\n        )\n\n        if chunk_id not in self._avg_cache:\n\n            chunk_path = os.path.join(\n                self.avg_dir,\n                f\"coco_avg_part_{chunk_id:03d}.pkl\"\n            )\n\n            with open(\n                chunk_path,\n                \"rb\"\n            ) as f:\n\n                self._avg_cache[\n                    chunk_id\n                ] = pickle.load(f)\n\n        return self._avg_cache[\n            chunk_id\n        ][image_id]\n\n\n    def __len__(self):\n\n        return len(\n            self.image_ids\n        )\n\n    def _load_image_from_disk(\n        self,\n        image_id\n    ):\n\n        path = self.image_paths[\n            image_id\n        ]\n\n        img = cv2.imread(\n            path\n        )\n\n        if img is None:\n\n            raise RuntimeError(\n                f\"Cannot read: {path}\"\n            )\n\n        return img\n\n    def _load_image(\n        self,\n        image_id\n    ):\n\n        return self._load_image_cached(\n            image_id\n        ).copy()\n\n\n    def _sample_object(\n        self,\n        image_id\n    ):\n\n        indices = (\n            self.objects_by_image[\n                image_id\n            ]\n        )\n\n        obj_idx = random.choice(\n            indices\n        )\n\n        obj = self.objects[\n            obj_idx\n        ]\n\n        return (\n            obj[\"bbox\"],\n            obj[\"category_id\"]\n        )\n\n\n    def _sample_hard_negative(\n        self,\n        template_category,\n        template_bbox,\n        template_image_id,\n    ):\n\n        cat_data = (\n            self._cat_arrays.get(\n                template_category\n            )\n        )\n\n        if (\n            cat_data is None\n            or len(cat_data[\"idxs\"]) <= 1\n        ):\n\n            return None\n\n        x, y, w, h = (\n            template_bbox\n        )\n\n        template_ratio = (\n            w / h\n        )\n\n        template_area = (\n            w * h\n        )\n\n        mask = (\n            cat_data[\"image_ids\"]\n            != template_image_id\n        )\n\n        if not mask.any():\n\n            return None\n\n        ratios = (\n            cat_data[\"ratios\"][mask]\n        )\n\n        areas = (\n            cat_data[\"areas\"][mask]\n        )\n\n        idxs = (\n            cat_data[\"idxs\"][mask]\n        )\n\n        ratio_diff = np.abs(\n            ratios\n            - template_ratio\n        )\n\n        size_diff = np.abs(\n            np.log(\n                areas\n                / template_area\n            )\n        )\n\n        score = (\n            ratio_diff\n            + size_diff\n        )\n\n        k = min(\n            50,\n            len(score)\n        )\n\n        top_k_idx = np.argpartition(\n            score,\n            k - 1\n        )[:k]\n\n        chosen = idxs[\n            np.random.choice(\n                top_k_idx\n            )\n        ]\n\n        return self.objects[\n            chosen\n        ]\n\n\n    def __getitem__(\n        self,\n        index\n    ):\n\n        template_image_id = (\n            self.image_ids[\n                index\n            ]\n        )\n\n        template_bbox, template_category = (\n            self._sample_object(\n                template_image_id\n            )\n        )\n\n        template_img = (\n            self._load_image(\n                template_image_id\n            )\n        )\n\n        template_avg = (\n            self._get_avg(\n                template_image_id\n            )\n        )\n\n\n        r = random.random()\n\n        hard_obj = None\n\n        if r < 1 / 3:\n\n\n            search_image_id = (\n                template_image_id\n            )\n\n            search_bbox = (\n                template_bbox\n            )\n\n            has_target = 1.0\n\n        elif r < 2 / 3:\n\n\n            neg_index = random.randrange(\n                len(self.image_ids)\n            )\n\n            search_image_id = (\n                self.image_ids[\n                    neg_index\n                ]\n            )\n\n            search_bbox, _ = (\n                self._sample_object(\n                    search_image_id\n                )\n            )\n\n            has_target = 0.0\n\n        else:\n\n\n            hard_obj = (\n                self._sample_hard_negative(\n                    template_category,\n                    template_bbox,\n                    template_image_id\n                )\n            )\n\n            if hard_obj is None:\n\n                neg_index = random.randrange(\n                    len(self.image_ids)\n                )\n\n                search_image_id = (\n                    self.image_ids[\n                        neg_index\n                    ]\n                )\n\n                search_bbox, _ = (\n                    self._sample_object(\n                        search_image_id\n                    )\n                )\n\n            else:\n\n                search_image_id = (\n                    hard_obj[\n                        \"image_id\"\n                    ]\n                )\n\n                search_bbox = (\n                    hard_obj[\n                        \"bbox\"\n                    ]\n                )\n\n            has_target = 0.0\n\n\n        search_img = (\n            self._load_image(\n                search_image_id\n            )\n        )\n\n        search_avg = (\n            self._get_avg(\n                search_image_id\n            )\n        )\n\n\n        sample = self.build_sample(\n\n            template_img=\n                template_img,\n\n            template_bbox=\n                template_bbox,\n\n            search_img=\n                search_img,\n\n            search_bbox=\n                search_bbox,\n\n            has_target=\n                has_target,\n\n            template_avg=\n                template_avg,\n\n            search_avg=\n                search_avg,\n        )\n\n        sample[\"dataset\"] = \"COCO\"\n\n        if has_target == 1.0:\n\n            sample[\"pair_type\"] = (\n                \"positive\"\n            )\n\n        elif hard_obj is not None:\n\n            sample[\"pair_type\"] = (\n                \"hard_negative\"\n            )\n\n        else:\n\n            sample[\"pair_type\"] = (\n                \"negative\"\n            )\n\n        return sample","metadata":{"_uuid":"9bd7397e-fd41-4f86-bebc-d50c0ec08089","_cell_guid":"6b21a06e-283e-44cb-a62c-b6ef0efd042b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-09-02T10:49:44.784388Z","iopub.execute_input":"2026-09-02T10:49:44.785278Z","iopub.status.idle":"2026-09-02T10:49:44.808173Z","shell.execute_reply.started":"2026-09-02T10:49:44.78525Z","shell.execute_reply":"2026-09-02T10:49:44.807138Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"got10k_dataset = GOT10kDataset(\n    got10k_dir=GOT10K_PATH,\n    metadata_path=(\n        \"/kaggle/working/\"\n        \"dasiam_preprocessed/\"\n        \"got10k/\"\n        \"got10k_metadata.pkl\"\n    ),\n    avg_dir=(\n        \"/kaggle/working/\"\n        \"dasiam_preprocessed/\"\n        \"got10k\"\n    ),\n    avg_chunk_size=500,\n    template_size=127,\n    search_size=271,\n    neg_pair_prob=0.2,\n)\n\ncoco_dataset = COCODataset(\n    coco_dir=COCO_PATH,\n    metadata_path=(\n        \"/kaggle/working/\"\n        \"dasiam_preprocessed/\"\n        \"coco/\"\n        \"coco_metadata.pkl\"\n    ),\n    avg_dir=(\n        \"/kaggle/working/\"\n        \"dasiam_preprocessed/\"\n        \"coco\"\n    ),\n    avg_chunk_size=5000,\n    template_size=127,\n    search_size=271,\n    neg_pair_prob=0.2,\n    hard_negative_prob=0.7,\n)\n\ntrain_dataset = ConcatDataset([got10k_dataset, coco_dataset])\nprint(len(train_dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:51:55.241043Z","iopub.execute_input":"2026-09-02T10:51:55.241437Z","iopub.status.idle":"2026-09-02T10:51:58.329208Z","shell.execute_reply.started":"2026-09-02T10:51:55.241399Z","shell.execute_reply":"2026-09-02T10:51:58.328425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport time\nfrom torch.utils.data import DataLoader\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=64,\n    shuffle=True,\n    num_workers=os.cpu_count(),\n    persistent_workers=True,\n    prefetch_factor=4,\n    pin_memory=True,\n    drop_last=True,\n)\n\nt0 = time.time()\nbatch = next(iter(train_loader))\nprint(\"full batch (данные):\", time.time() - t0)\n\ndata_iter = iter(train_loader)\nnext(data_iter)  # прогрев, не считаем\n\nt0 = time.time()\nbatch = next(data_iter)\nprint(\"batch (steady-state):\", time.time() - t0)\n\nimport time\n\n# GOT10k\nt0 = time.time()\nsample = got10k_dataset[0]\nprint(\"GOT10k getitem:\", time.time() - t0)\n\n# COCO\nt0 = time.time()\nsample = coco_dataset[0]\nprint(\"COCO getitem:\", time.time() - t0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:52:02.212293Z","iopub.execute_input":"2026-09-02T10:52:02.212639Z","iopub.status.idle":"2026-09-02T10:52:47.459094Z","shell.execute_reply.started":"2026-09-02T10:52:02.212612Z","shell.execute_reply":"2026-09-02T10:52:47.45494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\n\nt0 = time.time()\nfor i in range(200):\n    _ = got10k_dataset[0]\nelapsed = time.time() - t0\nprint(f\"200 сэмплов: {elapsed:.2f}с, среднее: {elapsed/200*1000:.1f}мс/сэмпл\")\nprint(\"размер avg-кэша:\", len(got10k_dataset._avg_cache))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:31:59.580454Z","iopub.execute_input":"2026-09-02T10:31:59.581152Z","iopub.status.idle":"2026-09-02T10:32:30.611156Z","shell.execute_reply.started":"2026-09-02T10:31:59.581097Z","shell.execute_reply":"2026-09-02T10:32:30.610349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\n\nt0 = time.time()\navg = np.mean(image, axis=(0, 1))\nprint(\"np.mean (full image):\", time.time() - t0)\n\nt0 = time.time()\ncrop = got10k_dataset.getSubwindow(image, (500, 500, 100, 100), 271, avg)\nprint(\"getSubwindow:\", time.time() - t0)\n\nt0 = time.time()\nresized = cv2.resize(crop, (127, 127))\nprint(\"resize:\", time.time() - t0)","metadata":{"trusted":true,"execution":{"execution_failed":"2026-08-25T22:13:46.182Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport random\n\n\ndef tensor_to_image(tensor):\n    \"\"\"\n    [C, H, W] -> [H, W, C]\n\n    Учитывает, что в _preprocess_image:\n        BGR -> RGB\n        image - avgChans\n    \"\"\"\n\n    img = tensor.detach().cpu().numpy()\n    img = np.transpose(img, (1, 2, 0))\n\n    img = img - img.min()\n\n    if img.max() > 0:\n        img = img / img.max()\n\n    return img\n\n\ndef draw_bbox(\n    image,\n    bbox,\n    color=None,\n    thickness=2,\n):\n    \"\"\"\n    bbox: [cx, cy, w, h]\n    \"\"\"\n\n    img = image.copy()\n\n    if color is None:\n        color = (1, 0, 0)\n\n    cx, cy, w, h = bbox\n\n    x1 = int(cx - w / 2)\n    y1 = int(cy - h / 2)\n    x2 = int(cx + w / 2)\n    y2 = int(cy + h / 2)\n\n    h_img, w_img = img.shape[:2]\n\n    x1 = max(0, min(w_img - 1, x1))\n    y1 = max(0, min(h_img - 1, y1))\n    x2 = max(0, min(w_img - 1, x2))\n    y2 = max(0, min(h_img - 1, y2))\n\n    cv2.rectangle(\n        img,\n        (x1, y1),\n        (x2, y2),\n        color,\n        thickness,\n    )\n\n    return img\n\n\ndef visualize_augmentation(\n    dataset,\n    index=None,\n):\n    \"\"\"\n    Показывает один sample после всех аугментаций.\n\n    Показывает:\n        1. Template\n        2. Search\n        3. Search + GT bbox\n    \"\"\"\n\n    if index is None:\n        index = random.randint(\n            0,\n            len(dataset) - 1\n        )\n\n    sample = dataset[index]\n\n    template = tensor_to_image(\n        sample[\"template\"]\n    )\n\n    search = tensor_to_image(\n        sample[\"search\"]\n    )\n\n    gt_bbox = sample[\"gt_bbox_patch\"].numpy()\n\n    search_bbox = draw_bbox(\n        search,\n        gt_bbox,\n    )\n\n    fig, axes = plt.subplots(\n        1,\n        3,\n        figsize=(15, 5),\n    )\n\n    axes[0].imshow(template)\n    axes[0].set_title(\n        \"Template\\n(after augmentation)\"\n    )\n    axes[0].axis(\"off\")\n\n    axes[1].imshow(search)\n    axes[1].set_title(\n        \"Search\\n(after augmentation)\"\n    )\n    axes[1].axis(\"off\")\n\n    axes[2].imshow(search_bbox)\n    axes[2].set_title(\n        f\"Search + GT bbox\\n\"\n        f\"cx={gt_bbox[0]:.1f}, \"\n        f\"cy={gt_bbox[1]:.1f}, \"\n        f\"w={gt_bbox[2]:.1f}, \"\n        f\"h={gt_bbox[3]:.1f}\"\n    )\n    axes[2].axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n\n\nfor _ in range(5):\n    visualize_augmentation(train_dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T10:32:55.669701Z","iopub.execute_input":"2026-09-02T10:32:55.670052Z","iopub.status.idle":"2026-09-02T10:32:57.864287Z","shell.execute_reply.started":"2026-09-02T10:32:55.67002Z","shell.execute_reply":"2026-09-02T10:32:57.863058Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import ConcatDataset\nimport random\n\n\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport random\n\n\nfound = {\n    \"positive\": None,\n    \"negative\": None,\n    \"hard_negative\": None\n}\n\n\nfor i in range(5000):\n\n    idx = random.randint(\n        0,\n        len(train_dataset)-1\n    )\n\n    sample = train_dataset[idx]\n\n\n    pair_type = sample.get(\n        \"pair_type\",\n        None\n    )\n\n\n    if pair_type in found:\n\n        if found[pair_type] is None:\n            found[pair_type] = sample\n\n\n    if all(\n        v is not None\n        for v in found.values()\n    ):\n        break\n\n\nfor pair_type, sample in found.items():\n\n    print(\"=\"*60)\n    print(pair_type)\n\n\n    if sample is None:\n        print(\"Не найден\")\n        continue\n\n\n    print(\n        \"Dataset:\",\n        sample[\"dataset\"]\n    )\n\n\n    print(\n        \"has_target:\",\n        sample[\"has_target\"].item()\n    )\n\n\n    template = sample[\"template\"].clone()\n    search = sample[\"search\"].clone()\n\n\n    template = template.permute(1,2,0).numpy()\n    search = search.permute(1,2,0).numpy()\n\n\n    template = (\n        template - template.min()\n    ) / (\n        template.max() - template.min()\n    )\n\n\n    search = (\n        search - search.min()\n    ) / (\n        search.max() - search.min()\n    )\n\n\n\n    plt.figure(figsize=(8,4))\n\n\n    plt.subplot(1,2,1)\n    plt.imshow(template)\n    plt.title(\n        f\"{pair_type}\\nTemplate\"\n    )\n    plt.axis(\"off\")\n\n\n    plt.subplot(1,2,2)\n    plt.imshow(search)\n    plt.title(\n        f\"{pair_type}\\nSearch\"\n    )\n    plt.axis(\"off\")\n\n\n    plt.show()","metadata":{"_uuid":"94b24dae-5642-4e1c-b95c-9539b6684c03","_cell_guid":"f54e3fb2-3595-4e52-8fd3-5a358cb99314","trusted":true,"collapsed":false,"execution":{"execution_failed":"2026-08-25T22:13:46.183Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=4,\n    shuffle=True,\n    num_workers=2,\n    drop_last=True\n)\n\ndata_iter = iter(train_loader)\nbatch = next(data_iter)\n\ntemplate = batch['template']       \nsearch = batch['search']           \ncls_targets = batch['cls_targets'] \nreg_targets = batch['reg_targets'] \n\nprint(\"=\" * 60)\nprint(\"УСПЕШНО ЗАГРУЖЕН БАТЧ ИЗ DATALOADER!\")\nprint(f\"Формат Template тензора: {template.shape}\")\nprint(f\"Формат Search тензора:   {search.shape}\")\nprint(f\"Формат Cls Targets:     {cls_targets.shape}\")\nprint(f\"Формат Reg Targets:     {reg_targets.shape}\")\nprint(\"=\" * 60)\n\ncls_map = cls_targets[0].numpy()\ncls_map = cls_map.reshape(5, 2, 19, 19)\n\nforeground = cls_map[:, 1].reshape(-1)\nbackground = cls_map[:, 0].reshape(-1)\n\npos_anchors = np.sum(foreground == 1)\nneg_anchors = np.sum(background == 1)\nignored_anchors = 1805 - pos_anchors - neg_anchors\n\nprint(\"СТАТИСТИКА АНКОРОВ ДЛЯ ПЕРВОГО КАДРА В БАТЧЕ:\")\nprint(f\" - Положительных (IoU >= 0.6): {pos_anchors}\")\nprint(f\" - Отрицательных  (IoU < 0.3):  {neg_anchors}\")\nprint(f\" - Игнорируемых               : {ignored_anchors}\")\nprint(f\" - Всего анкоров (5x19x19)   : {pos_anchors + neg_anchors + ignored_anchors}\")\nprint(\"=\" * 60)\n\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\n\n# 1. Извлекаем 1-й элемент из батча\nimg_template = template[0].permute(1, 2, 0).numpy() \nimg_search = search[0].permute(1, 2, 0).numpy()     \ngt_box = batch['gt_bbox_patch'][0].numpy()          \nhas_target = batch['has_target'][0].item()\n\ndef normalize_img(img):\n    img = img - img.min()\n    return img / (img.max() + 1e-5)\n\nimg_template = normalize_img(img_template)\nimg_search = normalize_img(img_search)\n\n\nif isinstance(train_dataset, torch.utils.data.ConcatDataset):\n    anchors = train_dataset.datasets[0].anchor_target_generator.anchors.numpy()\nelse:\n    anchors = train_dataset.anchor_target_generator.anchors.numpy()\ncls_map = cls_targets[0].numpy()\n\ncls_map = cls_map.reshape(5, 2, 19, 19)\n\nforeground = cls_map[:, 1, :, :]  \n\nforeground = foreground.reshape(-1)\n\npos_indices = np.where(foreground == 1)[0]\n\nfig, ax = plt.subplots(1, 2, figsize=(12, 6))\n\nax[0].imshow(img_template)\nax[0].set_title(\"Template (127x127)\")\nax[0].axis('off')\n\nax[1].imshow(img_search)\nis_pos_label = \"Positive Pair\" if has_target > 0.5 else \"Negative Pair (No Object)\"\nax[1].set_title(f\"Search Region (271x271) [{is_pos_label}]\")\n\nif has_target > 0.5:\n    for idx in pos_indices:\n        anc_cx, anc_cy, anc_w, anc_h = anchors[idx]\n        anc_x1 = anc_cx - anc_w / 2.0\n        anc_y1 = anc_cy - anc_h / 2.0\n        \n        rect_anc = patches.Rectangle(\n            (anc_x1, anc_y1), anc_w, anc_h,\n            linewidth=1, edgecolor='green', facecolor='none', alpha=0.5,\n            label='Positive Anchors' if idx == pos_indices[0] else \"\"\n        )\n        ax[1].add_patch(rect_anc)\n\n    gt_cx, gt_cy, gt_w, gt_h = gt_box\n    gt_x1 = gt_cx - gt_w / 2.0\n    gt_y1 = gt_cy - gt_h / 2.0\n\n    rect_gt = patches.Rectangle(\n        (gt_x1, gt_y1), gt_w, gt_h,\n        linewidth=2.5, edgecolor='red', facecolor='none',\n        label='Ground Truth GT'\n    )\n    ax[1].add_patch(rect_gt)\n    ax[1].legend(loc='upper right')\n\nax[1].axis('off')\nplt.tight_layout()\nplt.show()","metadata":{"_uuid":"6aed355b-84b3-4708-9346-2040cf121af9","_cell_guid":"5085f38d-e4db-4e04-b13f-a9c26bbfe4a4","trusted":true,"collapsed":false,"execution":{"execution_failed":"2026-08-25T22:13:46.183Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass SiamRPNLoss(nn.Module):\n    def __init__(self, rpn_weight: float = 1.0):\n        super().__init__()\n        self.rpn_weight = rpn_weight\n\n    def forward(self, pred_score, pred_delta, cls_targets, reg_targets, reg_weights):\n        \"\"\"\n        Args:\n            pred_score:  [B, 10, 19, 19] - логиты классификации от модели\n            pred_delta:  [B, 20, 19, 19] - сдвиги от модели\n            cls_targets: [B, 10, 19, 19] - бинарные маски (1.0 для BG/FG, 0.0 для ignore)\n            reg_targets: [B, 20, 19, 19] - целевые сдвиги\n            reg_weights: [B, 20, 19, 19] - маска позитивных анкоров (1.0 где pos, иначе 0.0)\n        \"\"\"\n        cls_mask = (cls_targets[:, 0::2, :, :] + cls_targets[:, 1::2, :, :]) > 0.0  # [B, 5, 19, 19]\n        cls_mask_2ch = cls_mask.repeat_interleave(2, dim=1)  # [B, 10, 19, 19] (bool)\n\n        num_cls = cls_mask_2ch.sum()\n\n        if num_cls > 0:\n            pred_score_pos = pred_score[cls_mask_2ch]\n            cls_targets_pos = cls_targets[cls_mask_2ch]\n\n            cls_loss = F.binary_cross_entropy_with_logits(\n                pred_score_pos, \n                cls_targets_pos, \n                reduction='mean'\n            )\n        else:\n            cls_loss = pred_score.sum() * 0.0  \n        pos_num = reg_weights.sum() / 4.0\n\n        pos_num_safe = torch.clamp(pos_num, min=1.0)\n\n        loss_reg_raw = F.smooth_l1_loss(\n            pred_delta, \n            reg_targets, \n            beta=1.0 / 9.0, \n            reduction='none'\n        )\n\n        loss_reg_masked = (loss_reg_raw * reg_weights).sum()\n\n        if pos_num > 0:\n            reg_loss = loss_reg_masked / pos_num_safe\n        else:\n            reg_loss = pred_delta.sum() * 0.0 \n\n        total_loss = cls_loss + self.rpn_weight * reg_loss\n\n        return total_loss, cls_loss, reg_loss","metadata":{"_uuid":"c62e045c-3af8-4267-8fb9-6285ffa5f5c0","_cell_guid":"c3bf869e-0aa5-43fd-9831-ba7200700d12","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-09-03T10:08:04.204466Z","iopub.execute_input":"2026-09-03T10:08:04.204689Z","iopub.status.idle":"2026-09-03T10:08:08.156024Z","shell.execute_reply.started":"2026-09-03T10:08:04.204666Z","shell.execute_reply":"2026-09-03T10:08:08.155212Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass DaSiamModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.Conv = nn.Conv2d(3, 96, kernel_size=11, stride=2, padding=0)\n        self.BatchNormalization = nn.BatchNorm2d(96)\n        self.pool = nn.MaxPool2d(kernel_size=3, stride=2)\n        self.relu = nn.ReLU()\n        \n        self.Conv_1 = nn.Conv2d(96, 256, kernel_size=5, stride=1, padding=0)\n        self.BatchNormalization_1 = nn.BatchNorm2d(256)\n        self.relu_1 = nn.ReLU()\n        self.pool_1 = nn.MaxPool2d(kernel_size=3, stride=2)\n        \n        self.Conv_2 = nn.Conv2d(256, 384, kernel_size=3, stride=1, padding=0)\n        self.BatchNormalization_2 = nn.BatchNorm2d(384)\n        self.relu_2 = nn.ReLU()\n        \n        self.Conv_3 = nn.Conv2d(384, 384, kernel_size=3, stride=1, padding=0)\n        self.BatchNormalization_3 = nn.BatchNorm2d(384)\n        self.relu_3 = nn.ReLU()\n        \n        self.Conv_4 = nn.Conv2d(384, 256, kernel_size=3, stride=1, padding=0)\n        self.BatchNormalization_4 = nn.BatchNorm2d(256)\n        \n        self.Conv_5 = nn.Conv2d(256, 256, kernel_size=3, stride=1, padding=0)\n        self.Conv_6 = nn.Conv2d(256, 20, kernel_size=4, stride=1, padding=0, bias=False)\n        self.Conv_7 = nn.Conv2d(20, 20, kernel_size=1, stride=1, padding=0)\n        \n        self.Conv_8 = nn.Conv2d(256, 256, kernel_size=3, stride=1, padding=0)\n        self.Conv_9 = nn.Conv2d(256, 10, kernel_size=4, stride=1, padding=0, bias=False)\n        \n        # Храним динамические веса в виде тензоров (без обрыва градиентов!)\n        self.r1_kernel = None\n        self.cls1_kernel = None\n\n    def set_adaptive_weights(self, r1_weights, cls1_weights):\n        \"\"\"\n        Сохраняем веса прямо как тензоры с градиентами.\n        r1_weights: [B, 20, 256, 4, 4] или [20, 256, 4, 4]\n        cls1_weights: [B, 10, 256, 4, 4] или [10, 256, 4, 4]\n        \"\"\"\n        self.r1_kernel = r1_weights\n        self.cls1_kernel = cls1_weights\n\n    def forward(self, x, use_adaptive=False):\n        x = self.relu(self.pool(self.BatchNormalization(self.Conv(x))))\n        x = self.relu_1(self.pool_1(self.BatchNormalization_1(self.Conv_1(x))))\n        x = self.relu_2(self.BatchNormalization_2(self.Conv_2(x)))\n        x = self.relu_3(self.BatchNormalization_3(self.Conv_3(x)))\n        x = self.BatchNormalization_4(self.Conv_4(x))\n        \n        intermediate_out = x\n        \n        if use_adaptive and self.r1_kernel is not None and self.cls1_kernel is not None:\n            x1 = self.Conv_5(x)  \n            \n            if self.r1_kernel.dim() == 5:\n                B, C, H, W = x1.shape\n                x1 = x1.view(1, B * C, H, W)\n                w_r1 = self.r1_kernel.view(B * 20, C, 4, 4)\n                x1 = F.conv2d(x1, w_r1, groups=B)\n                x1 = x1.view(B, 20, x1.shape[2], x1.shape[3])\n            else:\n                x1 = F.conv2d(x1, self.r1_kernel)\n\n            x1 = self.Conv_7(x1)\n\n            x2 = self.Conv_8(x) \n            \n            if self.cls1_kernel.dim() == 5:\n                B, C, H, W = x2.shape\n                x2 = x2.view(1, B * C, H, W)\n                w_cls1 = self.cls1_kernel.view(B * 10, C, 4, 4)\n                x2 = F.conv2d(x2, w_cls1, groups=B)\n                x2 = x2.view(B, 10, x2.shape[2], x2.shape[3])\n            else:\n                x2 = F.conv2d(x2, self.cls1_kernel)\n        else:\n            x1 = self.Conv_5(x)\n            x1 = self.Conv_6(x1)\n            x1 = self.Conv_7(x1)\n\n            x2 = self.Conv_8(x)\n            x2 = self.Conv_9(x2)\n        \n        return x1, x2, intermediate_out\n\n\nclass DaSiam_cls(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.Conv = nn.Conv2d(256, 2560, kernel_size=3, stride=1, padding=0)\n\n    def forward(self, x):\n        out = self.Conv(x)\n        B = out.shape[0]\n        return out.view(B, 10, 256, 4, 4)\n\n\nclass DaSiam_r(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.Conv = nn.Conv2d(256, 5120, kernel_size=3, stride=1, padding=0)\n\n    def forward(self, x):\n       \n        out = self.Conv(x)\n        B = out.shape[0]\n       \n        return out.view(B, 20, 256, 4, 4)","metadata":{"_uuid":"62dd22a6-e395-4ca3-b8bd-e45c723f62e4","_cell_guid":"8b23fa80-12b5-4b6b-8c02-c939c3b020ed","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-09-03T10:18:59.987344Z","iopub.execute_input":"2026-09-03T10:18:59.988159Z","iopub.status.idle":"2026-09-03T10:19:00.004181Z","shell.execute_reply.started":"2026-09-03T10:18:59.988127Z","shell.execute_reply":"2026-09-03T10:19:00.003313Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport time\nimport itertools\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\n\n\ndef save_dasiamrpn_weights(\n    model,\n    kernel_cls1,\n    kernel_r1,\n    save_dir=\"./weights\"\n):\n    os.makedirs(save_dir, exist_ok=True)\n\n    torch.save(\n        model.state_dict(),\n        os.path.join(save_dir, \"dasiamrpn_model.pth\")\n    )\n\n    torch.save(\n        kernel_cls1.state_dict(),\n        os.path.join(save_dir, \"dasiamrpn_kernel_cls1.pth\")\n    )\n\n    torch.save(\n        kernel_r1.state_dict(),\n        os.path.join(save_dir, \"dasiamrpn_kernel_r1.pth\")\n    )\n\n    print(f\"--> Веса успешно сохранены в каталог: {save_dir}\")\n\n\ntorch.backends.cudnn.benchmark = True\n\n\ndef train_dasiamrpn():\n\n    device = torch.device(\n        \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    )\n\n    print(f\"Используемое устройство: {device}\")\n\n    if device.type == \"cuda\":\n        print(\n            f\"GPU: {torch.cuda.get_device_name(0)}\"\n        )\n\n\n    model = DaSiamModel().to(device)\n    kernel_cls1 = DaSiam_cls().to(device)\n    kernel_r1 = DaSiam_r().to(device)\n\n\n    criterion = SiamRPNLoss(\n        rpn_weight=1.0\n    )\n\n\n    optimizer = torch.optim.SGD(\n        list(model.parameters())\n        + list(kernel_cls1.parameters())\n        + list(kernel_r1.parameters()),\n\n        lr=0.005,\n        momentum=0.9,\n        weight_decay=0.0005\n    )\n\n\n    scaler = torch.cuda.amp.GradScaler()\n\n\n\n\n\n    train_loader = DataLoader(\n        train_dataset,\n\n        batch_size=64,\n        shuffle=True,\n\n        num_workers=4,\n\n        pin_memory=True,\n        persistent_workers=True,\n        prefetch_factor=2,\n\n        drop_last=True\n    )\n\n\n    print(\n        f\"Размер датасета: {len(train_dataset)}\"\n    )\n\n    print(\n        f\"Количество batch'ей: {len(train_loader)}\"\n    )\n\n    print(\n        f\"Batch size: {train_loader.batch_size}\"\n    )\n\n    print(\n        f\"Workers: {train_loader.num_workers}\"\n    )\n\n\n    num_epochs = 40\n\n    best_loss = float(\"inf\")\n\n\n    for epoch in range(num_epochs):\n\n        model.train()\n        kernel_cls1.train()\n        kernel_r1.train()\n\n        running_loss = 0.0\n\n\n        for i, batch in enumerate(train_loader):\n\n            t0 = time.perf_counter()\n\n            template = batch[\"template\"]\n            search = batch[\"search\"]\n\n            cls_targets = batch[\"cls_targets\"]\n            reg_targets = batch[\"reg_targets\"]\n            reg_weights = batch[\"reg_weights\"]\n\n            t1 = time.perf_counter()\n\n\n            template = template.to(\n                device,\n                non_blocking=True\n            )\n\n            search = search.to(\n                device,\n                non_blocking=True\n            )\n\n            cls_targets = cls_targets.to(\n                device,\n                non_blocking=True\n            )\n\n            reg_targets = reg_targets.to(\n                device,\n                non_blocking=True\n            )\n\n            reg_weights = reg_weights.to(\n                device,\n                non_blocking=True\n            )\n\n\n            if device.type == \"cuda\":\n                torch.cuda.synchronize()\n\n            t2 = time.perf_counter()\n\n\n            optimizer.zero_grad(\n                set_to_none=True\n            )\n\n\n            with torch.amp.autocast(\"cuda\"):\n\n                _, _, template_feat = model(\n                    template,\n                    use_adaptive=False\n                )\n\n\n                cls1_weights = kernel_cls1(\n                    template_feat\n                )\n\n                r1_weights = kernel_r1(\n                    template_feat\n                )\n\n\n                B = template.size(0)\n\n\n                model.set_adaptive_weights(\n                    r1_weights.view(\n                        B,\n                        20,\n                        256,\n                        4,\n                        4\n                    ),\n\n                    cls1_weights.view(\n                        B,\n                        10,\n                        256,\n                        4,\n                        4\n                    )\n                )\n\n\n                pred_delta, pred_score, _ = model(\n                    search,\n                    use_adaptive=True\n                )\n\n\n                total_loss, cls_loss, reg_loss = criterion(\n                    pred_score,\n                    pred_delta,\n                    cls_targets,\n                    reg_targets,\n                    reg_weights\n                )\n\n\n            if device.type == \"cuda\":\n                torch.cuda.synchronize()\n\n            t3 = time.perf_counter()\n\n\n            scaler.scale(\n                total_loss\n            ).backward()\n\n\n            if device.type == \"cuda\":\n                torch.cuda.synchronize()\n\n            t4 = time.perf_counter()\n\n\n            scaler.unscale_(\n                optimizer\n            )\n\n\n            torch.nn.utils.clip_grad_norm_(\n                itertools.chain(\n                    model.parameters(),\n                    kernel_cls1.parameters(),\n                    kernel_r1.parameters()\n                ),\n                max_norm=10.0\n            )\n\n\n            scaler.step(\n                optimizer\n            )\n\n            scaler.update()\n\n\n            if device.type == \"cuda\":\n                torch.cuda.synchronize()\n\n            t5 = time.perf_counter()\n\n\n            running_loss += total_loss.item()\n\n\n            if (i + 1) % 10 == 0:\n\n                load_time = t1 - t0\n                transfer_time = t2 - t1\n                forward_time = t3 - t2\n                backward_time = t4 - t3\n                optimizer_time = t5 - t4\n\n                total_time = (\n                    load_time\n                    + transfer_time\n                    + forward_time\n                    + backward_time\n                    + optimizer_time\n                )\n\n\n                print(\n                    f\"\\n\"\n                    f\"Epoch [{epoch+1}/{num_epochs}] \"\n                    f\"Step [{i+1}/{len(train_loader)}]\\n\"\n\n                    f\"  Data preparation : \"\n                    f\"{load_time:.4f}s\\n\"\n\n                    f\"  CPU -> GPU       : \"\n                    f\"{transfer_time:.4f}s\\n\"\n\n                    f\"  Forward          : \"\n                    f\"{forward_time:.4f}s\\n\"\n\n                    f\"  Backward         : \"\n                    f\"{backward_time:.4f}s\\n\"\n\n                    f\"  Optimizer        : \"\n                    f\"{optimizer_time:.4f}s\\n\"\n\n                    f\"  TOTAL            : \"\n                    f\"{total_time:.4f}s\\n\"\n\n                    f\"  Loss             : \"\n                    f\"{total_loss.item():.4f} \"\n                    f\"(Cls: {cls_loss.item():.4f}, \"\n                    f\"Reg: {reg_loss.item():.4f})\"\n                )\n\n\n            elif (i + 1) % 50 == 0:\n\n                print(\n                    f\"Эпоха [{epoch+1}/{num_epochs}] | \"\n                    f\"Шаг [{i+1}/{len(train_loader)}] | \"\n                    f\"Loss: {total_loss.item():.4f} \"\n                    f\"(Cls: {cls_loss.item():.4f}, \"\n                    f\"Reg: {reg_loss.item():.4f})\"\n                )\n\n\n        epoch_loss = (\n            running_loss\n            / len(train_loader)\n        )\n\n\n        print(\n            f\"\\n\"\n            f\"========================================\\n\"\n            f\"Итог эпохи {epoch+1}\\n\"\n            f\"Средний Loss = {epoch_loss:.4f}\\n\"\n            f\"========================================\\n\"\n        )\n\n\n        if epoch_loss < best_loss:\n\n            best_loss = epoch_loss\n\n            save_dasiamrpn_weights(\n                model,\n                kernel_cls1,\n                kernel_r1,\n                save_dir=\"./best_weights\"\n            )\n","metadata":{"_uuid":"840809cf-a186-46f6-94b1-f81d112d6e16","_cell_guid":"c241412a-8fc4-4ba0-ad24-e09be8441368","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-09-03T10:30:11.057108Z","iopub.execute_input":"2026-09-03T10:30:11.057837Z","iopub.status.idle":"2026-09-03T10:30:11.077718Z","shell.execute_reply.started":"2026-09-03T10:30:11.057788Z","shell.execute_reply":"2026-09-03T10:30:11.076916Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dasiamrpn()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport cv2 \nimport numpy as np \nfrom typing import Tuple, List, Optional\nimport warnings\nimport torch\nimport torch.nn as nn\nimport cv2 \nimport numpy as np \nfrom typing import Tuple, List, Optional\nimport warnings\n# Определение класса трекера\nclass DaSiamRPNTracker:\n    def __init__(self, model_path: str, kernel_cls1_path: str, kernel_r1_path: str, device='cuda'):\n        # Проверяем доступность GPU\n        self.device = device\n        \n        # Флаги для контроля использования адаптивных весов\n        self.use_adaptive_weights = True\n        \n        # Загрузка моделей PyTorch\n        try:\n            self.siamRPN = DaSiamModel()\n            self.siamKernelCL1 = DaSiam_cls()\n            self.siamKernelR1 = DaSiam_r()\n\n            self._load_weights(model_path, kernel_cls1_path, kernel_r1_path)\n            self.siamRPN.eval()\n            self.siamKernelCL1.eval()\n            self.siamKernelR1.eval()\n\n            self.siamRPN.to(self.device)\n            self.siamKernelCL1.to(self.device)\n            self.siamKernelR1.to(self.device)\n            \n        except Exception as e:\n            raise ValueError(f\"Не удалось загрузить модели: {e}\")\n        \n        # Конфигурация трекера\n        self.trackState = {\n            'windowInfluence': 0.43,\n            'lr': 0.1,\n            'scale': 8,\n            'swapRB': True,  \n            'totalStride': 8,\n            'penaltyK': 0.055,\n            'exemplarSize': 127,\n            'instanceSize': 271,\n            'contextAmount': 0.5,\n            'ratios': [0.33, 0.5, 1.0, 2.0, 3.0],\n            'anchorNum': 5,\n            'anchors': None,\n            'windows': None,\n            'avgChans': None,\n            'imgSize': (0, 0),\n            'targetBox': None,\n            'scoreSize': 19,\n            'tracking_score': 0.0\n        }\n        \n        # Хранилище для шаблона\n        self.template_features = None\n        self.template_bbox = None\n        #self.verification_threshold = 0.8  # Порог для верификации\n        \n        self.is_initialized = False\n    \n    def _load_weights(self, model_path: str, kernel_cls1_path: str, kernel_r1_path: str) -> None:\n        \"\"\"Загрузка весов для моделей\"\"\"\n        try:\n            # Загружаем веса основной модели\n            if model_path:\n                checkpoint = torch.load(model_path, map_location=torch.device(self.device))\n                if 'state_dict' in checkpoint:\n                    self.siamRPN.load_state_dict(checkpoint['state_dict'])\n                else:\n                    self.siamRPN.load_state_dict(checkpoint)\n            \n            # Загружаем веса для kernel_cls1\n            if kernel_cls1_path:\n                checkpoint = torch.load(kernel_cls1_path, map_location=self.device)\n                if 'state_dict' in checkpoint:\n                    self.siamKernelCL1.load_state_dict(checkpoint['state_dict'])\n                else:\n                    self.siamKernelCL1.load_state_dict(checkpoint)\n            \n            # Загружаем веса для kernel_r1\n            if kernel_r1_path:\n                checkpoint = torch.load(kernel_r1_path, map_location=self.device)\n                if 'state_dict' in checkpoint:\n                    self.siamKernelR1.load_state_dict(checkpoint['state_dict'])\n                else:\n                    self.siamKernelR1.load_state_dict(checkpoint)\n                    \n        except Exception as e:\n            raise ValueError(f\"Ошибка загрузки весов: {e}\")\n    \n    def _preprocess_image(self, image: np.ndarray, target_size: Tuple[int, int]) -> torch.Tensor:\n        \"\"\"Препроцессинг изображения для PyTorch\"\"\"\n        # Изменение размера\n        if image.shape[:2] != target_size:\n            image = cv2.resize(image, target_size)\n        \n        # Конвертация из BGR в RGB если нужно\n        if len(image.shape) == 3 and image.shape[2] == 3:\n            if self.trackState['swapRB']:\n                image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        # Нормализация (DaSiamRPN использует другую нормализацию)\n        # Попробуем без нормализации или с простой нормализацией\n        image = image.astype(np.float32)\n        \n        # Попробуем вычесть среднее значение каналов\n        if self.trackState['avgChans'] is not None:\n            image = image - self.trackState['avgChans']\n        \n        # Конвертация в тензор и добавление batch dimension\n        image_tensor = torch.from_numpy(image).permute(2, 0, 1).unsqueeze(0).float()\n        \n        # Перемещение на устройство\n        image_tensor = image_tensor.to(self.device)\n        \n        return image_tensor\n    \n    def softmax(self, src: np.ndarray) -> np.ndarray:\n        \"\"\"Функция softmax\"\"\"\n        max_val = np.maximum(src[0], src[1])\n        exp_src = np.exp(src - max_val)\n        sum_val = exp_src[0] + exp_src[1] + 1e-8\n        dst = exp_src / sum_val\n        return dst\n    \n    def elementMax(self, src: np.ndarray) -> None:\n        \"\"\"Элементарный максимум\"\"\"\n        mask = src < 1.0\n        src[mask] = 1.0 / src[mask]\n    \n    def generateHanningWindow(self) -> np.ndarray:\n        \"\"\"Генерация окна Ханнинга\"\"\"\n        base_window = cv2.createHanningWindow((self.trackState['scoreSize'], \n                                             self.trackState['scoreSize']), cv2.CV_32F)\n        base_window = base_window.reshape(1, self.trackState['scoreSize'], self.trackState['scoreSize'])\n        hanning_windows = np.repeat(base_window, self.trackState['anchorNum'], axis=0)\n        return hanning_windows\n    \n    def generateAnchors(self) -> np.ndarray:\n        \"\"\"Генерация анкоров\"\"\"\n        totalStride = self.trackState['totalStride']\n        scales = self.trackState['scale']\n        scoreSize = self.trackState['scoreSize']\n        ratios = self.trackState['ratios']\n        anchorNum = self.trackState['anchorNum']\n        \n        base_anchors = []\n        for ratio in ratios:\n            size = totalStride * totalStride\n            ws = np.sqrt(size / ratio)\n            hs = ws * ratio\n            base_anchors.append([ws * scales, hs * scales])\n        \n        base_anchors = np.array(base_anchors, dtype=np.float32)\n        \n        anchors = np.zeros((4, anchorNum, scoreSize, scoreSize), dtype=np.float32)\n        ori = - (scoreSize // 2) * totalStride\n        \n        for i in range(scoreSize):\n            for j in range(scoreSize):\n                for k in range(anchorNum):\n                    anchors[0, k, i, j] = ori + totalStride * j\n                    anchors[1, k, i, j] = ori + totalStride * i\n                    anchors[2, k, i, j] = base_anchors[k, 0]\n                    anchors[3, k, i, j] = base_anchors[k, 1]\n        \n        return anchors\n    \n    def getSubwindow(self, img: np.ndarray, targetBox: Tuple[float, float, float, float], \n                    originalSize: float, avgChans: Tuple[float, float, float]) -> np.ndarray:\n        \"\"\"Извлечение подокна из изображения\"\"\"\n        img_h, img_w = img.shape[:2]\n        center_x, center_y, width, height = targetBox\n        c = (originalSize + 1) / 2\n        \n        xMin = int(np.round(center_x - c))\n        xMax = xMin + int(originalSize)\n        yMin = int(np.round(center_y - c))\n        yMax = yMin + int(originalSize)\n        \n        leftPad = max(0, -xMin)\n        topPad = max(0, -yMin)\n        rightPad = max(0, xMax - img_w)\n        bottomPad = max(0, yMax - img_h)\n        \n        xMin += leftPad\n        xMax += leftPad\n        yMin += topPad\n        yMax += topPad\n        \n        if topPad == 0 and bottomPad == 0 and leftPad == 0 and rightPad == 0:\n            crop = img[yMin:yMax, xMin:xMax].copy()\n        else:\n            padded_img = cv2.copyMakeBorder(img, topPad, bottomPad, leftPad, rightPad, \n                                         cv2.BORDER_CONSTANT, value=avgChans)\n            crop = padded_img[yMin:yMax, xMin:xMax].copy()\n        \n        return crop\n    \n    def init(self, image: np.ndarray, boundingBox: Tuple[int, int, int, int]) -> None:\n        \"\"\"Инициализация трекера\"\"\"\n        self.image_ = image.copy()\n        x, y, w, h = boundingBox\n        self.trackState['targetBox'] = (x + w/2, y + h/2, w, h)\n        self.trackState['imgSize'] = (image.shape[1], image.shape[0])\n        self.trackState['avgChans'] = np.mean(image, axis=(0, 1))\n        self.trackState['anchors'] = self.generateAnchors()\n        self.trackState['windows'] = self.generateHanningWindow()\n        \n        center_x, center_y, width, height = self.trackState['targetBox']\n        wc = width + self.trackState['contextAmount'] * (width + height)\n        hc = height + self.trackState['contextAmount'] * (width + height)\n        sz = np.sqrt(wc * hc)\n        \n        zCrop = self.getSubwindow(image, self.trackState['targetBox'], sz, self.trackState['avgChans'])\n        \n        # Препроцессинг для PyTorch\n        zCrop_tensor = self._preprocess_image(zCrop, (self.trackState['exemplarSize'], \n                                                      self.trackState['exemplarSize']))\n        \n        # Прямой проход через основную модель для получения шаблона\n        with torch.no_grad():\n            _, _, out1 = self.siamRPN(zCrop_tensor)\n\n            self.template_features = out1\n            self.template_bbox = boundingBox\n\n            cls1 = self.siamKernelCL1(out1)\n            r1 = self.siamKernelR1(out1)\n            \n            # Преобразуем kernel features в нужный формат\n            try:\n                if len(cls1.shape) == 4:  # [batch, channels, height, width]\n                    # Убираем batch dimension\n                    cls1 = cls1.squeeze(0)\n                    r1 = r1.squeeze(0)\n                \n                # Reshape к нужным размерам\n                if cls1.numel() == 10 * 256 * 4 * 4:\n                    cls1_reshaped = cls1.view(10, 256, 4, 4)\n                else:\n                    # Автоматический reshape\n                    cls1_reshaped = cls1.reshape(10, 256, 4, 4)\n                \n                if r1.numel() == 20 * 256 * 4 * 4:\n                    r1_reshaped = r1.view(20, 256, 4, 4)\n                else:\n                    # Автоматический reshape\n                    r1_reshaped = r1.reshape(20, 256, 4, 4)\n                self.siamRPN.set_adaptive_weights(r1_reshaped, cls1_reshaped)\n                self.use_adaptive_weights = True\n                \n            except Exception as e:\n                print(f\"Warning: Failed to set adaptive weights: {e}\")\n                print(\"Using default model weights\")\n                self.use_adaptive_weights = False\n            \n        self.is_initialized = True\n    \n    def update(self, image: np.ndarray) -> Tuple[Tuple[int, int, int, int], float]:\n        \"\"\"Обновление позиции трекера\"\"\"\n        if not self.is_initialized:\n            raise RuntimeError(\"Tracker not initialized. Call init() first.\")\n        \n        # self.image_ = image.copy()\n        # self.trackerEval(self.image_)\n        self.trackerEval(image)\n\n        \n        center_x, center_y, width, height = self.trackState['targetBox']\n        bbox = (int(center_x - width/2), int(center_y - height/2), int(width), int(height))\n        \n        return bbox, self.trackState['tracking_score']\n    \n    def trackerEval(self, img: np.ndarray) -> None:\n        \"\"\"Внутренняя оценка позиции\"\"\"\n        targetBox = self.trackState['targetBox']\n        center_x, center_y, width, height = targetBox\n        \n        wc = width + self.trackState['contextAmount'] * (width + height)\n        hc = height + self.trackState['contextAmount'] * (width + height)\n        sz = np.sqrt(wc * hc)\n        \n\n        #Вычисляем зум/масштаб, чтобы привести объект к размеру эталона (127x127)\n        scaleZ = self.trackState['exemplarSize'] / sz\n        \n        #Считаем, какую ширину кадра нужно вырезать вокруг объекта\n        searchSize = (self.trackState['instanceSize'] - self.trackState['exemplarSize']) / 2\n        pad = searchSize / scaleZ\n        sx = sz + 2 * pad\n        \n        xCrop = self.getSubwindow(img, targetBox, sx, self.trackState['avgChans'])\n        \n        # Препроцессинг для PyTorch\n        xCrop_tensor = self._preprocess_image(xCrop, (self.trackState['instanceSize'], \n                                                      self.trackState['instanceSize']))\n        \n        # Прямой проход через основную модель\n        with torch.no_grad():\n            delta_output, score_output, _ = self.siamRPN(xCrop_tensor, use_adaptive=self.use_adaptive_weights)\n            \n            # Конвертируем в numpy\n            delta = delta_output.cpu().numpy()\n            score = score_output.cpu().numpy()\n\n        # Убираем batch dimension\n        delta = delta.squeeze(0)  # [20, 19, 19]\n        score = score.squeeze(0)  # [10, 19, 19]\n        \n        try:\n            delta = delta.reshape(self.trackState['anchorNum'], 4,\n                       self.trackState['scoreSize'], self.trackState['scoreSize'])\n            delta = delta.transpose(1, 0, 2, 3)\n\n            score = score.reshape(self.trackState['anchorNum'], 2,\n                                self.trackState['scoreSize'], self.trackState['scoreSize'])\n            score = score.transpose(1, 0, 2, 3) \n\n            \n        except Exception as e:\n            print(\"Ошибка ресайза\")\n        \n        score = self.softmax(score)\n        score_obj = score[1].copy()\n        raw_idx = np.argmax(score_obj)\n        raw_anchor, raw_row, raw_col = np.unravel_index(raw_idx, score_obj.shape)\n  \n        anchors = self.trackState['anchors']\n        \n        delta[0] = delta[0] * anchors[2] + anchors[0]\n        delta[1] = delta[1] * anchors[3] + anchors[1]\n        delta[2] = np.exp(delta[2]) * anchors[2]\n        delta[3] = np.exp(delta[3]) * anchors[3]\n        \n        scaled_width = width * scaleZ\n        scaled_height = height * scaleZ\n        \n        sc = self.sizeCal(delta[2], delta[3]) / self.sizeCal(scaled_width, scaled_height)\n        self.elementMax(sc)\n        \n        rc = delta[2] / delta[3]\n        rc = (scaled_width / scaled_height) / rc\n        self.elementMax(rc)\n        \n        penalty = np.exp(-self.trackState['penaltyK'] * (rc * sc - 1))\n        \n        pscore = penalty * score_obj\n        pscore = pscore * (1.0 - self.trackState['windowInfluence']) + \\\n                self.trackState['windows'] * self.trackState['windowInfluence']\n        \n        best_idx = np.argmax(pscore)\n        best_score = pscore.flat[best_idx]\n        best_anchor, best_row, best_col = np.unravel_index(best_idx, pscore.shape)\n        \n        delta_flat = delta.reshape(4, -1)\n        best_delta = delta_flat[:, best_idx]\n        \n        res_x = best_delta[0] / scaleZ\n        res_y = best_delta[1] / scaleZ\n        res_width = best_delta[2] / scaleZ\n        res_height = best_delta[3] / scaleZ\n        \n        lr = penalty.flat[best_idx] * best_score * self.trackState['lr']\n        \n        new_center_x = center_x + res_x\n        new_center_y = center_y + res_y\n        new_width = width * (1 - lr) + res_width * lr\n        new_height = height * (1 - lr) + res_height * lr\n        \n        img_w, img_h = self.trackState['imgSize']\n        new_center_x = np.clip(new_center_x, 0, img_w)\n        new_center_y = np.clip(new_center_y, 0, img_h)\n        new_width = np.clip(new_width, 10, img_w)\n        new_height = np.clip(new_height, 10, img_h)\n        \n        self.trackState['targetBox'] = (new_center_x, new_center_y, new_width, new_height)\n        self.trackState['tracking_score'] = best_score\n    \n    def sizeCal(self, w: np.ndarray, h: np.ndarray) -> np.ndarray:\n        \"\"\"Расчет размера\"\"\"\n        pad = (w + h) * 0.5\n        sz2 = (w + pad) * (h + pad)\n        return np.sqrt(sz2 + 1e-8)\n    \n    def verify_target(self, image: np.ndarray, bbox: Tuple[int, int, int, int]) -> Tuple[bool, float]:\n        \"\"\"\n        Верификация цели с использованием DaSiamRPN\n        \n        Args:\n            image: Текущее изображение\n            bbox: Предполагаемое положение цели\n            \n        Returns:\n            Tuple[is_same_object, confidence_score]: \n            - is_same_object: True если это тот же объект\n            - confidence_score: Оценка уверенности (0-1)\n        \"\"\"\n        if not self.is_initialized:\n            raise RuntimeError(\"Трекер не инициализирован\")\n        \n        # Сохраняем текущее состояние\n        original_target_box = self.trackState['targetBox']\n        original_img_size = self.trackState['imgSize']\n        original_avg_chans = self.trackState['avgChans']\n        \n        try:\n            # Временно устанавливаем новую цель для верификации\n            x, y, w, h = bbox\n            self.trackState['targetBox'] = (x + w/2, y + h/2, w, h)\n            self.trackState['imgSize'] = (image.shape[1], image.shape[0])\n            self.trackState['avgChans'] = np.mean(image, axis=(0, 1))\n            \n            # Выполняем оценку для получения confidence score\n            self.trackerEval(image)\n            \n            # Используем confidence score от трекера как меру сходства\n            confidence = self.trackState['tracking_score']\n            is_same_object = confidence > 0.7\n            return is_same_object, confidence\n            \n        finally:\n            # Восстанавливаем исходное состояние\n            self.trackState['targetBox'] = original_target_box\n            self.trackState['imgSize'] = original_img_size\n            self.trackState['avgChans'] = original_avg_chans","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-03T10:08:47.924612Z","iopub.execute_input":"2026-09-03T10:08:47.925047Z","iopub.status.idle":"2026-09-03T10:08:48.154036Z","shell.execute_reply.started":"2026-09-03T10:08:47.925017Z","shell.execute_reply":"2026-09-03T10:08:48.153118Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\ndef compute_iou(boxA, boxB):\n    \"\"\"\n    box = [x, y, w, h]\n    \"\"\"\n\n    xA = max(boxA[0], boxB[0])\n    yA = max(boxA[1], boxB[1])\n\n    xB = min(\n        boxA[0] + boxA[2],\n        boxB[0] + boxB[2]\n    )\n\n    yB = min(\n        boxA[1] + boxA[3],\n        boxB[1] + boxB[3]\n    )\n\n    inter_w = max(0.0, xB - xA)\n    inter_h = max(0.0, yB - yA)\n\n    intersection = inter_w * inter_h\n\n    areaA = max(0.0, boxA[2]) * max(0.0, boxA[3])\n    areaB = max(0.0, boxB[2]) * max(0.0, boxB[3])\n\n    union = areaA + areaB - intersection\n\n    return intersection / (union + 1e-8)\n\n\n\ndef center_error(boxA, boxB):\n\n    centerA = (\n        boxA[0] + boxA[2] / 2,\n        boxA[1] + boxA[3] / 2\n    )\n\n    centerB = (\n        boxB[0] + boxB[2] / 2,\n        boxB[1] + boxB[3] / 2\n    )\n\n    return float(\n        np.hypot(\n            centerA[0] - centerB[0],\n            centerA[1] - centerB[1]\n        )\n    )\n\n\ndef evaluate_dasiam_tracker(\n    tracker,\n    video_folders,\n    max_videos=None\n):\n\n    if max_videos is not None:\n        videos = video_folders[:max_videos]\n    else:\n        videos = video_folders\n\n    all_ious = []\n    all_center_errors = []\n    all_norm_errors = []\n\n    video_results = []\n\n    print()\n    print(\"=\" * 70)\n    print(\"START EVALUATION\")\n    print(\"=\" * 70)\n    print(f\"Videos: {len(videos)}\")\n    print(\"=\" * 70)\n\n    for video_dir in tqdm(\n        videos,\n        desc=\"Evaluating videos\"\n    ):\n\n        video_name = os.path.basename(video_dir)\n\n        gt_path = os.path.join(\n            video_dir,\n            \"groundtruth.txt\"\n        )\n\n        if not os.path.exists(gt_path):\n\n            print(\n                f\"[WARNING] No groundtruth: {video_name}\"\n            )\n\n            continue\n\n\n        gt_boxes = []\n\n        with open(\n            gt_path,\n            \"r\"\n        ) as f:\n\n            for line in f:\n\n                line = line.strip()\n\n                if not line:\n                    continue\n\n                line = line.replace(\n                    \"\\t\",\n                    \",\"\n                )\n\n                values = [\n                    float(v)\n                    for v in line.split(\",\")\n                ]\n\n                gt_boxes.append(values)\n\n        if len(gt_boxes) == 0:\n            continue\n\n\n        image_paths = [\n            os.path.join(\n                video_dir,\n                f\"{i + 1:08d}.jpg\"\n            )\n            for i in range(\n                len(gt_boxes)\n            )\n        ]\n\n\n        frame0 = cv2.imread(\n            image_paths[0]\n        )\n\n        if frame0 is None:\n\n            print(\n                f\"[WARNING] Cannot read first frame: \"\n                f\"{video_name}\"\n            )\n\n            continue\n\n\n        tracker.init(\n            frame0,\n            tuple(gt_boxes[0])\n        )\n\n        video_ious = []\n        video_errors = []\n        video_norm_errors = []\n\n        video_ious.append(1.0)\n        video_errors.append(0.0)\n        video_norm_errors.append(0.0)\n\n        for frame_idx in range(\n            1,\n            len(gt_boxes)\n        ):\n\n            frame = cv2.imread(\n                image_paths[frame_idx]\n            )\n\n            if frame is None:\n                continue\n\n            try:\n\n                pred_box, score = tracker.update(\n                    frame\n                )\n\n            except Exception as e:\n\n                print(\n                    f\"\\n[ERROR] {video_name}, \"\n                    f\"frame {frame_idx}: {e}\"\n                )\n\n                break\n\n            gt_box = gt_boxes[frame_idx]\n\n\n            iou = compute_iou(\n                pred_box,\n                gt_box\n            )\n\n\n            ce = center_error(\n                pred_box,\n                gt_box\n            )\n\n\n            gt_w = gt_box[2]\n            gt_h = gt_box[3]\n\n            gt_diag = np.hypot(\n                gt_w,\n                gt_h\n            ) + 1e-8\n\n            normalized_error = (\n                ce / gt_diag\n            )\n\n            video_ious.append(iou)\n            video_errors.append(ce)\n            video_norm_errors.append(\n                normalized_error\n            )\n\n\n        video_ious_np = np.asarray(\n            video_ious,\n            dtype=np.float32\n        )\n\n        video_errors_np = np.asarray(\n            video_errors,\n            dtype=np.float32\n        )\n\n        video_norm_errors_np = np.asarray(\n            video_norm_errors,\n            dtype=np.float32\n        )\n\n        video_ao = np.mean(\n            video_ious_np\n        )\n\n        video_sr50 = np.mean(\n            video_ious_np > 0.5\n        )\n\n        video_sr75 = np.mean(\n            video_ious_np > 0.75\n        )\n\n        video_precision20 = np.mean(\n            video_errors_np < 20\n        )\n\n        video_norm_precision = np.mean(\n            video_norm_errors_np < 0.2\n        )\n\n        video_results.append(\n            {\n                \"video\": video_name,\n                \"frames\": len(video_ious_np),\n                \"AO\": video_ao,\n                \"SR@0.5\": video_sr50,\n                \"SR@0.75\": video_sr75,\n                \"Precision@20px\": video_precision20,\n                \"NormalizedPrecision@0.2\":\n                    video_norm_precision\n            }\n        )\n\n        all_ious.extend(\n            video_ious\n        )\n\n        all_center_errors.extend(\n            video_errors\n        )\n\n        all_norm_errors.extend(\n            video_norm_errors\n        )\n\n    return (\n        np.asarray(\n            all_ious,\n            dtype=np.float32\n        ),\n        np.asarray(\n            all_center_errors,\n            dtype=np.float32\n        ),\n        np.asarray(\n            all_norm_errors,\n            dtype=np.float32\n        ),\n        pd.DataFrame(\n            video_results\n        )\n    )\n\n\n\ndef calculate_metrics(\n    ious,\n    center_errors,\n    norm_errors\n):\n\n\n    ao = np.mean(\n        ious\n    )\n\n\n    sr50 = np.mean(\n        ious > 0.5\n    )\n\n    sr75 = np.mean(\n        ious > 0.75\n    )\n\n    thresholds = np.linspace(\n        0,\n        1,\n        101\n    )\n\n    success_rates = np.array([\n        np.mean(\n            ious >= threshold\n        )\n        for threshold in thresholds\n    ])\n\n    success_auc = np.trapezoid(\n        success_rates,\n        thresholds\n    )\n\n\n    precision20 = np.mean(\n        center_errors <= 20\n    )\n\n\n\n    normalized_precision = np.mean(\n        norm_errors <= 0.2\n    )\n\n    return {\n        \"AO\": ao,\n        \"SR@0.5\": sr50,\n        \"SR@0.75\": sr75,\n        \"Success AUC\": success_auc,\n        \"Precision@20px\": precision20,\n        \"Normalized Precision@0.2\":\n            normalized_precision\n    }\n\n\ndef plot_success_curve(ious):\n\n    thresholds = np.linspace(\n        0,\n        1,\n        101\n    )\n\n    success_rates = np.array([\n        np.mean(\n            ious >= threshold\n        )\n        for threshold in thresholds\n    ])\n\n    auc = np.trapezoid(\n        success_rates,\n        thresholds\n    )\n\n    plt.figure(\n        figsize=(7, 5)\n    )\n\n    plt.plot(\n        thresholds,\n        success_rates\n    )\n\n    plt.xlabel(\n        \"IoU threshold\"\n    )\n\n    plt.ylabel(\n        \"Success rate\"\n    )\n\n    plt.title(\n        f\"Success Plot — AUC = {auc:.4f}\"\n    )\n\n    plt.xlim(\n        0,\n        1\n    )\n\n    plt.ylim(\n        0,\n        1\n    )\n\n    plt.grid(\n        True\n    )\n\n    plt.show()\n\n    return auc\n\n\ndef plot_precision_curve(\n    center_errors\n):\n\n    thresholds = np.arange(\n        0,\n        51,\n        1\n    )\n\n    precision = np.array([\n        np.mean(\n            center_errors <= threshold\n        )\n        for threshold in thresholds\n    ])\n\n    precision20 = np.mean(\n        center_errors <= 20\n    )\n\n    plt.figure(\n        figsize=(7, 5)\n    )\n\n    plt.plot(\n        thresholds,\n        precision\n    )\n\n    plt.xlabel(\n        \"Center error (pixels)\"\n    )\n\n    plt.ylabel(\n        \"Precision\"\n    )\n\n    plt.title(\n        f\"Precision Plot — P@20 = {precision20:.4f}\"\n    )\n\n    plt.xlim(\n        0,\n        50\n    )\n\n    plt.ylim(\n        0,\n        1\n    )\n\n    plt.grid(\n        True\n    )\n\n    plt.show()\n\n    return precision20\n\n\ndef plot_normalized_precision(\n    norm_errors\n):\n\n    thresholds = np.linspace(\n        0,\n        0.5,\n        101\n    )\n\n    precision = np.array([\n        np.mean(\n            norm_errors <= threshold\n        )\n        for threshold in thresholds\n    ])\n\n    precision20 = np.mean(\n        norm_errors <= 0.2\n    )\n\n    plt.figure(\n        figsize=(7, 5)\n    )\n\n    plt.plot(\n        thresholds,\n        precision\n    )\n\n    plt.xlabel(\n        \"Normalized center error\"\n    )\n\n    plt.ylabel(\n        \"Precision\"\n    )\n\n    plt.title(\n        f\"Normalized Precision — P@0.2 = {precision20:.4f}\"\n    )\n\n    plt.xlim(\n        0,\n        0.5\n    )\n\n    plt.ylim(\n        0,\n        1\n    )\n\n    plt.grid(\n        True\n    )\n\n    plt.show()\n\n    return precision20\n\n\n\ndevice = torch.device(\n    \"cuda\"\n    if torch.cuda.is_available()\n    else \"cpu\"\n)\n\nprint(\n    f\"Device: {device}\"\n)\n\nif torch.cuda.is_available():\n\n    print(\n        f\"GPU: {torch.cuda.get_device_name(0)}\"\n    )\n\n\n\ntracker = DaSiamRPNTracker(\n\n    model_path=\n        \"/kaggle/input/datasets/viollett/dasiamrpn-files/dasiamrpn_model.pth\",\n\n    kernel_cls1_path=\n        \"/kaggle/input/datasets/viollett/dasiamrpn-files/dasiamrpn_kernel_cls1.pth\",\n\n    kernel_r1_path=\n        \"/kaggle/input/datasets/viollett/dasiamrpn-files/dasiamrpn_kernel_r1.pth\",\n\n    device=device\n)\n\nprint(\n    \"Tracker loaded successfully.\"\n)\n\n\nval_dir = os.path.join(\n    GOT10K_PATH,\n    \"val\"\n)\n\nval_video_folders = sorted(\n    os.path.join(\n        val_dir,\n        f\n    )\n    for f in os.listdir(\n        val_dir\n    )\n    if os.path.isdir(\n        os.path.join(\n            val_dir,\n            f\n        )\n    )\n)\n\nprint(\n    f\"GOT-10k validation videos: \"\n    f\"{len(val_video_folders)}\"\n)\n\n\n\nMAX_VIDEOS = 180\n\n\nious, center_errors, norm_errors, video_results = \\\n    evaluate_dasiam_tracker(\n        tracker,\n        val_video_folders,\n        max_videos=MAX_VIDEOS\n    )\n\n\n\nmetrics = calculate_metrics(\n    ious,\n    center_errors,\n    norm_errors\n)\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL GOT-10k RESULTS\")\nprint(\"=\" * 70)\n\nfor name, value in metrics.items():\n\n    print(\n        f\"{name:30s}: {value:.4f}\"\n    )\n\nprint(\"=\" * 70)\n\n\nsuccess_auc = plot_success_curve(\n    ious\n)\n\nprecision20 = plot_precision_curve(\n    center_errors\n)\n\nnormalized_precision = \\\n    plot_normalized_precision(\n        norm_errors\n    )\n\n\n\nRESULTS_DIR = (\n    \"/kaggle/working/\"\n    \"dasiam_evaluation\"\n)\n\nos.makedirs(\n    RESULTS_DIR,\n    exist_ok=True\n)\n\n\nmetrics_df = pd.DataFrame(\n    [\n        metrics\n    ]\n)\n\nmetrics_path = os.path.join(\n    RESULTS_DIR,\n    \"got10k_metrics.csv\"\n)\n\nmetrics_df.to_csv(\n    metrics_path,\n    index=False\n)\n\n\n\n\nvideo_results_path = os.path.join(\n    RESULTS_DIR,\n    \"got10k_per_video.csv\"\n)\n\nvideo_results.to_csv(\n    video_results_path,\n    index=False\n)\n\n\nnp.save(\n    os.path.join(\n        RESULTS_DIR,\n        \"ious.npy\"\n    ),\n    ious\n)\n\n\nnp.save(\n    os.path.join(\n        RESULTS_DIR,\n        \"center_errors.npy\"\n    ),\n    center_errors\n)\n\n\n# Normalized errors\n\nnp.save(\n    os.path.join(\n        RESULTS_DIR,\n        \"normalized_errors.npy\"\n    ),\n    norm_errors\n)\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"RESULTS SAVED\")\nprint(\"=\" * 70)\n\nprint(\n    metrics_path\n)\n\nprint(\n    video_results_path\n)\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"WORST 10 VIDEOS\")\nprint(\"=\" * 70)\n\nworst = video_results.sort_values(\n    \"AO\",\n    ascending=True\n).head(10)\n\nprint(\n    worst.to_string(\n        index=False\n    )\n)\n\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"BEST 10 VIDEOS\")\nprint(\"=\" * 70)\n\nbest = video_results.sort_values(\n    \"AO\",\n    ascending=False\n).head(10)\n\nprint(\n    best.to_string(\n        index=False\n    )\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-03T10:36:21.95752Z","iopub.execute_input":"2026-09-03T10:36:21.958395Z","iopub.status.idle":"2026-09-03T10:46:09.149598Z","shell.execute_reply.started":"2026-09-03T10:36:21.958363Z","shell.execute_reply":"2026-09-03T10:46:09.149003Z"}},"outputs":[],"execution_count":null}]}