{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"datasetVersion","sourceId":15433411,"datasetId":9872944,"databundleVersionId":16352653}],"dockerImageVersionId":31328,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===============================\n# FINAL WORKING VERSION (COO MASK VERSION)\n# ===============================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\n# ===== 1. Check dataset =====\nprint(\"===== DATA STRUCTURE =====\")\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    print(dirname)\n\n# ===== 2. Correct path =====\nBASE_DIR = '/kaggle/input/datasets/tylerde/ecg-training-data-0001-only/data'\n\nIMG_DIR = os.path.join(BASE_DIR, 'rectified')\nMASK_DIR = os.path.join(BASE_DIR, 'masks')\n\n# ===== 3. Build dataframe =====\nimage_files = sorted(os.listdir(IMG_DIR))\nmask_files = sorted(os.listdir(MASK_DIR))\n\n# 对齐 id（去掉后缀）\ndef get_id_from_img(x):\n    return x.split('-')[0]\n\ndef get_id_from_mask(x):\n    return x.split('.')[0]\n\nimg_map = {get_id_from_img(f): f for f in image_files}\nmask_map = {get_id_from_mask(f): f for f in mask_files}\n\ncommon_ids = sorted(list(set(img_map.keys()) & set(mask_map.keys())))\n\ndf = pd.DataFrame({\n    'id': common_ids\n})\n\nprint(\"Samples:\", len(df))\n\n# ===== 4. Load COO mask =====\ndef load_coo_mask(npz_path):\n    data = np.load(npz_path)\n\n    C, H, W = data['shape']\n    mask = np.zeros((C, H, W), dtype=np.float32)\n\n    for i in range(C):\n        ys = data[f'ch{i}_y']\n        xs = data[f'ch{i}_x']\n        vs = data[f'ch{i}_v']\n        mask[i, ys, xs] = vs\n\n    return mask\n\n# ===== 5. Dataset =====\nclass ECGDataset(Dataset):\n    def __init__(self, df, img_dir, mask_dir, img_size=512):\n        self.df = df\n        self.img_dir = img_dir\n        self.mask_dir = mask_dir\n        self.img_size = img_size\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        sid = self.df.iloc[idx]['id']\n\n        img_path = os.path.join(self.img_dir, f\"{sid}-0001.rect.png\")\n        mask_path = os.path.join(self.mask_dir, f\"{sid}.mask-coo.npz\")\n\n        # ===== image =====\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = img.astype(np.float32) / 255.0\n\n        # ===== mask (COO → dense) =====\n        mask = load_coo_mask(mask_path)\n\n        # 👉 只取第一个 channel（或你也可以做 multi-channel）\n        mask = mask[0]\n\n        # ===== resize =====\n        img = cv2.resize(img, (self.img_size, self.img_size))\n        mask = cv2.resize(mask, (self.img_size, self.img_size))\n\n        # ===== format =====\n        img = np.transpose(img, (2, 0, 1))\n        mask = np.expand_dims(mask, axis=0)\n\n        return torch.tensor(img, dtype=torch.float32), torch.tensor(mask, dtype=torch.float32)\n\n# ===== 6. DataLoader =====\ndataset = ECGDataset(df, IMG_DIR, MASK_DIR)\n\ntrain_loader = DataLoader(\n    dataset,\n    batch_size=2,\n    shuffle=True,\n    num_workers=2\n)\n\n# ===== 7. Model =====\nclass PatchEmbedding(nn.Module):\n    def __init__(self, img_size=512, patch_size=16, embed_dim=256):\n        super().__init__()\n        self.proj = nn.Conv2d(3, embed_dim, patch_size, patch_size)\n\n    def forward(self, x):\n        x = self.proj(x)\n        x = x.flatten(2).transpose(1, 2)\n        return x\n\nclass SimpleViT(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.patch = PatchEmbedding()\n\n        num_patches = (512 // 16) ** 2\n        self.pos = nn.Parameter(torch.randn(1, num_patches, 256))\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=256, nhead=8, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, 4)\n\n        self.decoder = nn.ConvTranspose2d(256, 1, 16, 16)\n\n    def forward(self, x):\n        x = self.patch(x)\n        x = x + self.pos\n\n        x = self.transformer(x)\n\n        B, N, C = x.shape\n        H = W = int(N ** 0.5)\n\n        x = x.permute(0, 2, 1).contiguous().view(B, C, H, W)\n        x = self.decoder(x)\n\n        return x\n\n# ===== 8. Train =====\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nmodel = SimpleViT().to(device)\n\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\nloss_fn = nn.BCEWithLogitsLoss()\n\nfor epoch in range(3):\n    model.train()\n    total_loss = 0\n\n    for images, masks in train_loader:\n        images = images.to(device)\n        masks = masks.to(device)\n\n        outputs = model(images)\n\n        loss = loss_fn(outputs, masks)\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n\n    print(f\"Epoch {epoch+1}, Loss: {total_loss/len(train_loader):.4f}\")\n\n# ===== 9. Save =====\ntorch.save(model.state_dict(), '/kaggle/working/transformer_ecg.pth')\n\nprint(\"DONE\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-10T04:00:22.827647Z","iopub.execute_input":"2026-04-10T04:00:22.827903Z","iopub.status.idle":"2026-04-10T05:03:06.411183Z","shell.execute_reply.started":"2026-04-10T04:00:22.827860Z","shell.execute_reply":"2026-04-10T05:03:06.409033Z"}},"outputs":[],"execution_count":null}]}