{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":14401911,"sourceType":"datasetVersion","datasetId":8937412},{"sourceId":14464474,"sourceType":"datasetVersion","datasetId":9234686},{"sourceId":14475589,"sourceType":"datasetVersion","datasetId":9135537},{"sourceId":14505899,"sourceType":"datasetVersion","datasetId":9260270},{"sourceId":14553044,"sourceType":"datasetVersion","datasetId":9246587},{"sourceId":14564526,"sourceType":"datasetVersion","datasetId":8837688},{"sourceId":14572023,"sourceType":"datasetVersion","datasetId":8784735},{"sourceId":14577204,"sourceType":"datasetVersion","datasetId":9311582},{"sourceId":703783,"sourceType":"modelInstanceVersion","modelInstanceId":523681,"modelId":537702}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"try:\n    import cc3d\nexcept:\n    #https://pypi.org/project/connected-components-3d/\n    #!pip install connected-components-3d\n\n    !ls /kaggle/input/hengck-libs/setup\n    !pip install connected-components-3d --no-index --find-links=file:///kaggle/input/hengck-libs/setup/\n\nimport cc3d\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nimport matplotlib\n#matplotlib.use('TkAgg')\nimport shutil\nimport os\nfrom tqdm import tqdm\nimport sys\nsys.path.append('/kaggle/input/hengck-libs')\n\nprint('import ok!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:55:16.067762Z","iopub.execute_input":"2026-01-20T07:55:16.068366Z","iopub.status.idle":"2026-01-20T07:55:16.075402Z","shell.execute_reply.started":"2026-01-20T07:55:16.068338Z","shell.execute_reply":"2026-01-20T07:55:16.074543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nMODE   = 'submit'  # submit  local fake\nDEVICE = 'cuda'\nFLOAT_TYPE = torch.float32 #torch.bfloat16\nFAIL_ID = []\n\nKAGGLE_DIR = \\\n\t'/kaggle/input/physionet-ecg-image-digitization'\nWEIGHT_DIR = \\\n\t'/kaggle/input/hengck-libs/weight'\nOUT_DIR = \\\n    f'/kaggle/working/output-{MODE}'\n\ndef make_test_fake_df(): \n    valid_df = pd.read_csv(f'{KAGGLE_DIR}/train.csv')\n    valid_df.loc[:,'id']=valid_df['id'].astype(str) \n    fake_test_df=[]\n    for i,d in valid_df.iterrows():\n        #if i==4: break\n        image_id = d['id']\n    \n        truth_df = pd.read_csv(f'{KAGGLE_DIR}/train/{image_id}/{image_id}.csv')\n        non_nan_count = truth_df.count()\n        #print(i,image_id,non_nan_count)\n        #print(non_nan_count.index)\n    \n        #lead\tfs\tnumber_of_rows \n        this_df = pd.DataFrame({\n            'id':image_id ,\n            'lead':non_nan_count.index,\n            'fs': d['fs'],\n            'number_of_rows':non_nan_count.values \n        })\n        fake_test_df.append(this_df)\n        if i==0: print(this_df)\n    fake_test_df = pd.concat(fake_test_df)\n    return fake_test_df\n\n\n# set valid/test data\nif MODE == 'local':\n    from sample_list import ERROR_ID\n    valid_df = pd.read_csv(f'{KAGGLE_DIR}/train.csv')\n    valid_df['id']=valid_df['id'].astype(str)\n    \n    valid_id = [\n        f'{image_id}-{type_id}' for image_id in ERROR_ID\n        #f'{image_id}-{type_id}' for image_id in valid_df['id'].values[500:]\n        for type_id in ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012']\n    ]\n    valid_id = [\n        '11842146-0012','144746082-0009','225208096-0006', '2289894144-0012','1617515072-0006',\n        '2289894144-0010','2566168201-0009', '2659677149-0011'\n    ]\n    \n    df_ = pd.read_csv('/kaggle/input/hengck-libs/valid.csv')\n    valid_id = []\n    for i in range(len(df_)):\n        cur_df = df_.loc[i]\n        n = cur_df['image_path'].split('/')[-1].split('.')[0]\n        valid_id.append(n)\n    valid_id = valid_id[:20]\n    \nif MODE == 'submit':\n\tvalid_df = pd.read_csv(f'{KAGGLE_DIR}/test.csv')\n\tvalid_df['id']=valid_df['id'].astype(str) \n\tvalid_id = valid_df['id'].unique().tolist()\n\nif MODE == 'fake':\n\tvalid_df = make_test_fake_df()\n\tvalid_df['id']=valid_df['id'].astype(str) \n\tvalid_id = valid_df['id'].unique().tolist()\n\n#--------------------------------------\n\ndef read_image(sample_id):\n    if MODE == 'local':\n        image_id, type_id = sample_id.split('-')\n        image = cv2.imread(f'{KAGGLE_DIR}/train/{image_id}/{image_id}-{type_id}.png', cv2.IMREAD_COLOR_RGB)\n        return image\n    if MODE == 'submit':\n        image_id = sample_id\n        image = cv2.imread(f'{KAGGLE_DIR}/test/{image_id}.png', cv2.IMREAD_COLOR_RGB)\n        return image\n    if MODE == 'fake':\n        image_id = sample_id \n        type_id = ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012'][\n            int(image_id)%9\n        ] \n        image = cv2.imread(f'{KAGGLE_DIR}/train/{image_id}/{image_id}-{type_id}.png', cv2.IMREAD_COLOR_RGB)\n        return image\n\ndef read_sampling_length(sample_id):\n\tif MODE == 'local':\n\t\timage_id, type_id = sample_id.split('-')\n\t\td = valid_df[valid_df['id']==image_id].iloc[0]\n\t\tlength = d.sig_len\n\t\treturn length\n\tif MODE == 'submit':\n\t\timage_id = sample_id\n\t\td = valid_df[\n\t\t\t(valid_df['id']==image_id) & (valid_df['lead']=='II')\n\t\t].iloc[0]\n\t\tlength = d.number_of_rows\n\t\treturn length\n\tif MODE == 'fake':\n\t\timage_id = sample_id\n\t\td = valid_df[\n\t\t\t(valid_df['id']==image_id) & (valid_df['lead']=='II')\n\t\t].iloc[0]\n\t\tlength = d.number_of_rows\n\t\treturn length\n\n#valid_id = valid_id[:300]\nprint('valid_id:', len(valid_id))\nprint('\\t', valid_id[:3], '...')\nprint('setting ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:55:16.076706Z","iopub.execute_input":"2026-01-20T07:55:16.076915Z","iopub.status.idle":"2026-01-20T07:55:16.170574Z","shell.execute_reply.started":"2026-01-20T07:55:16.076895Z","shell.execute_reply":"2026-01-20T07:55:16.169841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage0\nprint('*** STARTING STAGE0 ***')\n\nfrom stage0_model import Net as Stage0Net\n#from stage0_common import *\nimport stage0_common\n\ndef change_color(image_rgb):\n    hsv = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2HSV)\n    h, s, v = cv2.split(hsv)\n    \n    v_denoised = cv2.fastNlMeansDenoising(v, h=5.5)\n    \n    std = np.std(v_denoised)\n    clip_limit = max(1.0, min(3.5, 2.0 + std / 25))\n    clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=(8, 8))\n    v_enhanced = clahe.apply(v_denoised)\n    \n    hsv_enhanced = cv2.merge([h, s, v_enhanced])\n    return cv2.cvtColor(hsv_enhanced, cv2.COLOR_HSV2RGB)\ndef run_stage0(stage0_net,sample_id):\n    image = read_image(sample_id)\n    image_for_model = change_color(image)\n    batch1 = stage0_common.image_to_batch(image_for_model)\n    batch2 = stage0_common.image_to_batch(image)\n    with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n        with torch.no_grad():\n            output = stage0_net(batch1)\n            output2 = stage0_net(batch2)\n            output['marker']  = (output['marker'] + output2['marker']) / 2.\n            output['orientation']  = (output['orientation'] + output2['orientation']) / 2.\n            rotated, keypoint = stage0_common.output_to_predict(image, batch1, output)\n            normalised, keypoint, homo = stage0_common.normalise_by_homography(rotated, keypoint)\n    torch.cuda.empty_cache()\n    return normalised\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:55:16.171395Z","iopub.execute_input":"2026-01-20T07:55:16.171658Z","iopub.status.idle":"2026-01-20T07:55:16.179519Z","shell.execute_reply.started":"2026-01-20T07:55:16.171636Z","shell.execute_reply":"2026-01-20T07:55:16.178903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage1\nprint('*** STARTING STAGE1 ***')\n\nfrom stage1_model import Net as Stage1Net\nimport stage1_common\nimport stage1_common_origin\n\ndef run_stage1(stage1_net,image):\n\n    batch = {\n        'image': torch.from_numpy(np.ascontiguousarray(image.transpose(2, 0, 1))).unsqueeze(0),\n    }\n    num_tta = 1\n    with torch.amp.autocast('cuda', dtype=FLOAT_TYPE): #torch.bfloat16\n        with torch.no_grad():\n            output = stage1_net(batch)\n            \n            gridpoint_xy_1, more = stage1_common.output_to_predict(image, batch, output)\n            rectified_1 = stage1_common.rectify_image(image, gridpoint_xy_1)\n            empty_image_1 = np.zeros((1700,2200, 3), np.uint8)\n            empty_image_1[:,:2166] = rectified_1\n\n            gridpoint_xy_2, more = stage1_common_origin.output_to_predict(image, batch, output)\n            rectified_2 = stage1_common_origin.rectify_image(image, gridpoint_xy_2)\n            \n    torch.cuda.empty_cache()\n    return empty_image_1,rectified_2\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:55:16.180498Z","iopub.execute_input":"2026-01-20T07:55:16.180738Z","iopub.status.idle":"2026-01-20T07:55:16.200148Z","shell.execute_reply.started":"2026-01-20T07:55:16.180717Z","shell.execute_reply":"2026-01-20T07:55:16.199379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage2\nimport torch.nn as nn\nimport torch.nn.functional as F\nclass DWSeparableConv1d(nn.Module):\n    def __init__(self, in_ch, out_ch, k=7, dilation=1, p_drop=0.1):\n        super().__init__()\n        self.k = k\n        self.d = dilation\n        self.dw = nn.Conv1d(in_ch, in_ch, k, padding=0, dilation=dilation, groups=in_ch, bias=False)\n        self.pw = nn.Conv1d(in_ch, out_ch, 1, bias=False)\n        self.gn = nn.GroupNorm(1, out_ch)\n        self.act = nn.SiLU()\n        self.drop = nn.Dropout(p_drop)\n\n    def forward(self, x):\n        pad = ((self.k - 1) // 2) * self.d\n        if pad > 0:\n            x = torch.cat([x[..., -pad:], x, x[..., :pad]], dim=-1)  # circular padding\n        x = self.dw(x)\n        x = self.pw(x)\n        x = self.gn(x)\n        x = self.act(x)\n        x = self.drop(x)\n        return x\n\nclass ResidualBlock(nn.Module):\n    def __init__(self, ch, k=7, dilation=1, p_drop=0.1):\n        super().__init__()\n        self.conv1 = DWSeparableConv1d(ch, ch, k=k, dilation=dilation, p_drop=p_drop)\n        self.conv2 = DWSeparableConv1d(ch, ch, k=k, dilation=dilation, p_drop=p_drop)\n\n    def forward(self, x):\n        out = self.conv1(x)\n        out = self.conv2(out)\n        return x + out\n\nclass ResidualLSTM(nn.Module):\n    def __init__(self, channels, hidden=128, num_layers=1, dropout=0.1, bidirectional=True):\n        super().__init__()\n        self.lstm = nn.LSTM(\n            input_size=channels,\n            hidden_size=hidden,\n            num_layers=num_layers,\n            batch_first=True,\n            bidirectional=bidirectional,\n            dropout=dropout if num_layers > 1 else 0.0,\n        )\n        out_ch = hidden * (2 if bidirectional else 1)\n        self.proj = nn.Sequential(nn.Linear(out_ch, channels), nn.SiLU())\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, x):\n        B, C, L = x.shape\n        h, _ = self.lstm(x.transpose(1, 2).contiguous())\n        h = self.proj(h).transpose(1, 2).contiguous()\n        return x + self.dropout(h)\n\n\n\nclass MultiScaleTCN(nn.Module):\n    \"\"\"\n    并行多尺度膨胀堆栈（ASPP风格）：每个分支一个 dilation 列表，最后通道融合。\n    \"\"\"\n    def __init__(self, ch, k=7, p_drop=0.1, branches=( (1,2,4), (1,3,9), (2,4,8,16) )):\n        super().__init__()\n        self.branches = nn.ModuleList()\n        for dil_list in branches:\n            blocks = nn.Sequential(*[ResidualBlock(ch, k=k, dilation=d, p_drop=p_drop) for d in dil_list])\n            self.branches.append(blocks)\n        self.fuse = nn.Sequential(\n            nn.Conv1d(len(branches) * ch, ch, kernel_size=1, bias=False),\n            nn.GroupNorm(1, ch),\n            nn.SiLU()\n        )\n\n    def forward(self, x):\n        outs = [b(x) for b in self.branches]\n        h = torch.cat(outs, dim=1)  # concat channels\n        h = self.fuse(h)\n        return h\n\nclass FourierBranch(nn.Module):\n    \"\"\"\n    分段多通道 Fourier 特征（严格均匀分段）\n    - 将序列按 n_segments 均分：seg_len = L // n_segments\n    - 仅对前 seg_len*n_segments 的部分做分段周期特征（严格均匀）\n    - 每段：每通道找主峰频率（去 DC），取 1..K 次谐波的幅值+相位\n    - 对每段得到一个 out_ch 向量，并广播回该段对应的时间范围\n    输出:\n      h: (B,out_ch,L)  （尾部 L - seg_len*n_segments 的位置填 0）\n      gate: (B,out_ch,L)\n    \"\"\"\n    def __init__(self, in_ch: int, k_harmonics: int = 6, out_ch: int = 128, p_drop: float = 0.1, n_segments: int = 4):\n        super().__init__()\n        self.in_ch = in_ch\n        self.k = k_harmonics\n        self.n_segments = n_segments\n        self.out_ch = out_ch\n\n        # 对每个 segment 单独映射到 out_ch\n        self.fc_seg = nn.Sequential(\n            nn.Linear(in_ch * 2 * k_harmonics, out_ch),\n            nn.SiLU(),\n            nn.Dropout(p_drop),\n        )\n\n        self.gate = nn.Sequential(\n            nn.Conv1d(out_ch, out_ch, 1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x: torch.Tensor):\n        # x: (B,C,L)\n        B, C, L = x.shape\n        seg_len = L // self.n_segments\n        L_use = seg_len * self.n_segments  # 严格均匀覆盖长度\n\n        # 输出在时间轴上铺开：每段一个向量 -> 该段内常量特征\n        h = x.new_zeros(B, self.out_ch, L)\n\n        if seg_len == 0:\n            gate = self.gate(h)\n            return h, gate\n\n        for s in range(self.n_segments):\n            a = s * seg_len\n            b = a + seg_len\n            xs = x[..., a:b]  # (B,C,seg_len)\n\n            # 分段 FFT\n            Xf = torch.fft.rfft(xs, dim=-1)           # (B,C,F)\n            mag = torch.abs(Xf)                       # (B,C,F)\n            mag_no_dc = mag[..., 1:]                  # (B,C,F-1)\n\n            peak_idx = torch.argmax(mag_no_dc, dim=-1) + 1  # (B,C)\n            Fmax = mag.size(-1)\n\n            feats = []\n            for k in range(1, self.k + 1):\n                idx = (peak_idx * k).clamp_max(Fmax - 1).unsqueeze(-1).long()  # (B,C,1)\n                amp_k = torch.gather(mag, -1, idx).squeeze(-1)                 # (B,C)\n                phase_k = torch.angle(torch.gather(Xf, -1, idx)).squeeze(-1)   # (B,C)\n                feats.append(amp_k)\n                feats.append(phase_k)\n\n            feat_seg = torch.cat(feats, dim=-1)       # (B, C*2K)\n            vec = self.fc_seg(feat_seg)               # (B, out_ch)\n            h[..., a:b] = vec.unsqueeze(-1).expand(-1, -1, seg_len)\n\n        gate = self.gate(h)\n        return h, gate\n\nclass ImputeTCNPeriodicMS(nn.Module):\n    \"\"\"\n    多尺度 + 周期谐波融合 + LSTM 残差\n    输入:\n      x_masked: (B, 4, L)\n      mask:     (B, 4, L)  mask=1表示缺失\n    输出:\n      y:        (B, 4, L)\n    \"\"\"\n    def __init__(self, in_ch: int = 4,\n                 hidden=128, k=7, p_drop=0.1,\n                 lstm_hidden=128, lstm_layers=1, lstm_dropout=0.1,\n                 k_harmonics=6,\n                 ms_branches=((1,2,4), (1,3,9), (2,4,8,16))):\n        super().__init__()\n        self.in_ch = in_ch\n\n        # 拼接 [x_masked, mask] 后通道变成 2*in_ch\n        self.stem = nn.Conv1d(2 * in_ch, hidden, kernel_size=1)\n        self.ms_tcn = MultiScaleTCN(ch=hidden, k=k, p_drop=p_drop, branches=ms_branches)\n\n        self.fourier = FourierBranch(\n            in_ch=in_ch,\n            k_harmonics=k_harmonics,\n            out_ch=hidden,\n            p_drop=p_drop,\n            n_segments=4\n        )\n\n        self.seq_block = ResidualLSTM(\n            channels=hidden,\n            hidden=lstm_hidden,\n            num_layers=lstm_layers,\n            dropout=lstm_dropout,\n            bidirectional=True\n        )\n\n        # 输出 4 通道\n        self.head = nn.Conv1d(hidden, in_ch, kernel_size=1)\n\n    def forward(self, x_masked, mask):\n        # x_masked, mask: (B,4,L), mask=1表示缺失\n        x_in = torch.cat([x_masked, mask], dim=1)  # (B,8,L)\n\n        h = self.stem(x_in)\n        h = self.ms_tcn(h)\n\n        # Fourier 分支不要直接用 x_masked（缺失段被置0会影响频谱）\n        # 用“观测均值填充缺失段”得到相对干净的频域输入\n        obs = 1.0 - mask\n        eps = 1e-6\n        mean_obs = (x_masked * obs).sum(dim=-1, keepdim=True) / obs.sum(dim=-1, keepdim=True).clamp_min(eps)\n        x_for_fft = x_masked * obs + mean_obs * mask\n\n        h_freq, gate = self.fourier(x_for_fft)\n        h = h * gate + h_freq\n\n        h = self.seq_block(h)\n\n        # 残差式：只学习缺失点应该加多少修正\n        delta = self.head(h)           # (B,4,L)\n        y = x_masked + mask * delta\n\n        # 保险：观测点严格保持输入\n        y = mask * y + (1.0 - mask) * x_masked\n        return y\nimport torchvision\nprint('*** STARTING STAGE2 ***')\ndef fill_miss_runs(series, miss_mask_,zero_mv, min_len):\n    \"\"\"\n    series, miss_mask: shape (4, length)\n    min_len: minimum consecutive True count in miss_mask to fill\n    \"\"\"\n    miss_mask = miss_mask_.copy()\n    series = series.copy()\n    for lead in range(series.shape[0]):\n        mask = miss_mask[lead]\n        nonzero = series[lead][(~mask) & (series[lead] != 0)]\n        fill_value = nonzero.mean() if nonzero.size > 0 else 0.0\n\n        edges = np.diff(np.concatenate(([0], mask.astype(np.int8), [0])))\n        starts = np.where(edges == 1)[0]\n        ends = np.where(edges == -1)[0]\n\n        for s, e in zip(starts, ends):\n            if e - s >= min_len:\n                series[lead, s:e] = fill_value\n                miss_mask[lead,s:e] = 0\n    return series,miss_mask\n\n#=======================\nfrom scipy.signal import savgol_filter, medfilt\nfrom scipy.signal import butter, filtfilt, find_peaks\ndef fill_missing_by_phase_within_each_1250(\n    ecg_4xT: np.ndarray,\n    mask_4xT: np.ndarray,\n    rpeaks: np.ndarray,\n    total_len: int = 5000,\n    seg_len: int = 1250,\n    leads: list[int] | None = None,        # None => all channels(0..3)\n    prefer_near_cycles: bool = True,       # 保留参数（这里不影响“取平均”策略）\n    search_radius: int = 2,                # src点缺失时，在src±radius内找最近非缺失点（同一1250段内）\n) -> np.ndarray:\n    \"\"\"\n    更新后的填充逻辑（符合你给的例子）：\n    - 在每个lead、每个1250段内处理\n    - 对缺失点 t：先找本段内最近R峰 p_ref，算 k = t - p_ref\n    - 对本段内其它R峰 p_j：src = p_j + k（例如 50 对应 100 的k=-50，则 600->550, 1100->1050）\n    - 如果多个周期（多个p_j）都能提供有效src，则把这些值取平均填回 t\n    - 只要求“src点”（或src附近±search_radius）不缺失，不要求整拍完整\n    \"\"\"\n    X = np.asarray(ecg_4xT, dtype=np.float32)\n    M = np.asarray(mask_4xT).astype(np.uint8)\n    r = np.sort(np.asarray(rpeaks, dtype=int))\n\n    if X.shape != M.shape:\n        raise ValueError(f\"mask shape {M.shape} must match ecg {X.shape}\")\n    if X.shape[1] != total_len:\n        raise ValueError(f\"Expected total_len={total_len}, got T={X.shape[1]}\")\n    if total_len % seg_len != 0:\n        raise ValueError(\"total_len must be divisible by seg_len\")\n\n    C, T = X.shape\n    nseg = total_len // seg_len\n\n    r = r[(r >= 0) & (r < T)]\n    if r.size < 2:\n        return X.copy()\n\n    if leads is None:\n        leads = list(range(C))\n\n    out = X.copy()\n\n    def _nearest_valid_in_segment(c: int, src: int, seg_a: int, seg_b: int) -> int | None:\n        \"\"\"在同一segment[seg_a, seg_b)内，从src开始向两侧找最近非缺失点\"\"\"\n        if not (seg_a <= src < seg_b):\n            return None\n        if M[c, src] == 0:\n            return int(src)\n        for d in range(1, max(0, search_radius) + 1):\n            for cand in (src - d, src + d):\n                if seg_a <= cand < seg_b and M[c, cand] == 0:\n                    return int(cand)\n        return None\n\n    for c in leads:\n        for s in range(nseg):\n            seg_a, seg_b = s * seg_len, (s + 1) * seg_len\n\n            # 本段内R峰\n            rs = r[(r >= seg_a) & (r < seg_b)]\n            if rs.size < 2:\n                continue\n\n            # 本段内缺失点\n            miss = np.flatnonzero(M[c, seg_a:seg_b] == 1)\n            if miss.size == 0:\n                continue\n            miss = miss + seg_a\n\n            for t in miss:\n                # 参考峰：本段内离t最近的R峰\n                ref_i = int(np.argmin(np.abs(rs - t)))\n                p_ref = int(rs[ref_i])\n                k = int(t) - p_ref\n\n                vals = []\n                for j in range(rs.size):\n                    if j == ref_i:\n                        continue\n                    p_j = int(rs[j])\n                    src = p_j + k\n\n                    # 只允许同一1250段内\n                    if not (seg_a <= src < seg_b):\n                        continue\n\n                    src_ok = _nearest_valid_in_segment(c, src, seg_a, seg_b)\n                    if src_ok is None:\n                        continue\n\n                    vals.append(float(X[c, src_ok]))\n\n                if len(vals) > 0:\n                    out[c, int(t)] = float(np.mean(vals))\n                # 否则：保持缺失不动\n\n    return out\n\ndef _bandpass(x, fs, lo=5.0, hi=18.0, order=3):\n    # 比较通用的QRS频段；若T波误检多，可把hi降到15\n    b, a = butter(order, [lo/(fs/2), hi/(fs/2)], btype=\"bandpass\")\n    return filtfilt(b, a, x)\n\ndef _mad(x, eps=1e-9):\n    med = np.median(x)\n    return np.median(np.abs(x - med)) + eps\n\ndef _moving_average(x, k):\n    if k <= 1:\n        return x\n    w = np.ones(k, dtype=np.float32) / k\n    return np.convolve(x, w, mode=\"same\")\n\ndef _pt_mwi(x_bp, fs):\n    dx = np.diff(x_bp, prepend=x_bp[0])\n    sq = dx * dx\n    win = max(1, int(0.15 * fs))  # 150ms\n    return _moving_average(sq, win)\n\ndef _searchback(peaks, score, fs, refractory, rr_ratio=1.66, relax=0.5):\n    # peaks: 初检峰（在score上），score: 能量包络\n    if len(peaks) < 2:\n        return peaks\n    rr = np.diff(peaks)\n    rr_med = np.median(rr)\n    med = np.median(score)\n    mad = _mad(score)\n    base_thr = med + 4.0 * mad\n    relax_thr = med + (4.0 * relax) * mad\n\n    out = [int(peaks[0])]\n    for i in range(1, len(peaks)):\n        prev = out[-1]\n        cur = int(peaks[i])\n        gap = cur - prev\n        if gap > rr_ratio * rr_med:\n            seg = score[prev:cur]\n            if len(seg) > refractory:\n                extra, _ = find_peaks(seg, distance=refractory, height=relax_thr)\n                if len(extra) > 0:\n                    best = int(extra[np.argmax(seg[extra])])\n                    out.append(prev + best)\n        out.append(cur)\n\n    out = np.array(sorted(set(out)), dtype=int)\n    return out\n\ndef _pick_reference_lead(X_bp):\n    \"\"\"\n    选一个用于精定位的导联：QRS能量更强者更适合。\n    X_bp: (4,T)\n    \"\"\"\n    # 用带通后微分能量作为评分\n    scores = []\n    for c in range(4):\n        dx = np.diff(X_bp[c], prepend=X_bp[c, 0])\n        scores.append(float(np.mean(dx*dx)))\n    return int(np.argmax(scores))\n\ndef _ensure_upright(x_bp, rough_peaks, fs, check_ms=120):\n    if len(rough_peaks) == 0:\n        return x_bp\n    w = int(fs * check_ms / 1000.0)\n    vals = []\n    for p in rough_peaks[:min(len(rough_peaks), 10)]:\n        s = max(0, p - w)\n        e = min(len(x_bp), p + w + 1)\n        seg = x_bp[s:e]\n        if len(seg) < 5:\n            continue\n        pos = float(np.max(seg))\n        neg = float(-np.min(seg))\n        vals.append(pos - neg)\n    if vals and np.median(vals) < 0:\n        return -x_bp\n    return x_bp\ndef _moving_stats(x, k):\n    \"\"\"\n    简单滑窗均值/方差（用卷积实现），k为窗口点数\n    返回 mean, std\n    \"\"\"\n    if k <= 1:\n        m = x\n        s = np.ones_like(x) * (np.std(x) + 1e-6)\n        return m, s\n    w = np.ones(k, dtype=np.float32) / k\n    mean = np.convolve(x, w, mode=\"same\")\n    mean2 = np.convolve(x*x, w, mode=\"same\")\n    var = np.maximum(mean2 - mean*mean, 1e-12)\n    std = np.sqrt(var)\n    return mean, std\n\ndef _local_zscore(x, fs, win_s=1.5):\n    k = max(3, int(fs * win_s))\n    mu, sd = _moving_stats(x, k)\n    return (x - mu) / (sd + 1e-6)\ndef detect_global_rpeaks(ecg_4xT, fs=500,\n                         refractory_ms=200,\n                         thr_k=4.0,\n                         fuse_lo=5.0, fuse_hi=18.0,\n                         refine_ms=80,\n                         min_bpm=40, max_bpm=200):\n    \"\"\"\n    返回：全局R峰索引（samples）\n    ecg_4xT: (4,T)\n    \"\"\"\n    X = np.asarray(ecg_4xT, dtype=np.float32)\n    if X.shape[0] != 4:\n        raise ValueError(f\"Expected (4,T), got {X.shape}\")\n\n    # 1) 每导联带通 + MWI\n    X_bp = np.zeros_like(X)\n    M = np.zeros_like(X)\n    for c in range(4):\n        x = X[c] - np.mean(X[c])\n        x_bp = _bandpass(x, fs, lo=fuse_lo, hi=fuse_hi, order=3)\n        X_bp[c] = x_bp\n        M[c] = _pt_mwi(x_bp, fs)\n\n    # 2) 融合能量包络（对形态波动鲁棒）\n    E = np.sqrt(np.mean(M, axis=0) + 1e-12)\n\n    # 3) 自适应阈值 + 初检\n    refractory = int(fs * refractory_ms / 1000.0)\n    Z = _local_zscore(E, fs, win_s=1.5)\n\n    refractory = int(fs * refractory_ms / 1000.0)\n    z_thr = 1.0\n    peaks, _ = find_peaks(Z, distance=refractory, height=z_thr)\n    peaks = _searchback(peaks.astype(int), Z, fs, refractory, rr_ratio=1.66, relax=0.6)\n    if len(peaks) == 0:\n        return np.array([], dtype=int)\n\n    # 4) 搜索回补漏\n    peaks = _searchback(peaks.astype(int), E, fs, refractory, rr_ratio=1.66, relax=0.5)\n\n    # 5) RR范围清理（全局）\n    min_rr = int(fs * 60.0 / max_bpm)\n    max_rr = int(fs * 60.0 / min_bpm)\n    if len(peaks) >= 2:\n        rr = np.diff(peaks)\n        good = np.ones(len(peaks), dtype=bool)\n        bad = np.where((rr < min_rr) | (rr > max_rr))[0] + 1\n        good[bad] = False\n        peaks = peaks[good]\n\n    # 6) 精定位到“正波峰”\n    ref = _pick_reference_lead(X_bp)\n    x_ref = _ensure_upright(X_bp[ref], peaks, fs, check_ms=120)\n\n    w = int(fs * refine_ms / 1000.0)\n    refined = []\n    for p in peaks:\n        s = max(0, p - w)\n        e = min(X.shape[1], p + w + 1)\n        seg = x_ref[s:e]\n        if len(seg) < 3:\n            refined.append(int(p))\n        else:\n            refined.append(int(s + np.argmax(seg)))  # 只取波峰\n    refined = np.array(refined, dtype=int)\n\n    # 最终去重（防止精定位把两个峰挤得太近）\n    if len(refined) >= 2:\n        refined = np.sort(refined)\n        keep = [refined[0]]\n        for p in refined[1:]:\n            if p - keep[-1] >= refractory:\n                keep.append(p)\n        refined = np.array(keep, dtype=int)\n\n    return refined\n\n\n\n#======================\n\n\n\n\n\n\ndef sliding_window_quarter_segmentation(\n    image,\n    models,\n    window_size: int = 528*1,\n    overlap_frac: float = 0.25,\n    overlap_px: int = 256*1,\n    interp_mode: str = \"bilinear\",\n    align_corners: bool = True,\n    device = 'cuda',\n    scalesize = 4,\n    crop_remain = 64\n    \n):\n\n    image = image.unsqueeze(0)\n    B, C, H, W = image.shape\n    # Compute stride from overlap\n    if overlap_px is not None:\n        stride = window_size - overlap_px\n    else:\n        stride = max(1, int(round(window_size * (1 - overlap_frac))))\n    stride = max(1, stride)\n\n    # Compute window start indices covering the width\n    starts = []\n    ends = []\n    s = 0\n    while True:\n        starts.append(s)\n        ends.append(s + window_size)\n        if s + window_size >= W:\n            break\n        s_next = s + stride\n        # Ensure last window touches the end\n        if s_next + window_size > W:\n            s = W - window_size\n        else:\n            s = s_next\n        if starts and s == starts[-1]:\n            break\n\n    image_up = F.interpolate(\n        image, size=(H, W * scalesize), mode=interp_mode, align_corners=align_corners\n    )\n\n    preds = []\n    with torch.no_grad():\n        for s in starts:\n            # 在放大后的图上切片：把原始 s/window_size 映射到放大尺度\n            s_up = s * scalesize\n            e_up = (s + window_size) * scalesize\n            slice_img = image_up[:, :, :, s_up:e_up]\n\n            inputs = {\"image\": slice_img}\n            pmy = None\n            for model in models:\n                pred = model(inputs)[\"pixel\"]\n                pmy = pred if pmy is None else (pmy + pred)\n            pmy /= len(models)\n            preds.append(pmy)\n\n    K = preds[0].shape[1]\n    out = torch.zeros((1, K, H, W * scalesize), dtype=preds[0].dtype, device=device)\n    mask = torch.zeros((1, 1, H, W * scalesize), dtype=preds[0].dtype, device=device)\n\n    for mm, (start, end) in enumerate(zip(starts, ends)):\n        start_up = start * scalesize\n        end_up = end * scalesize\n        cr = crop_remain * scalesize\n\n        if mm == 0:\n            out[:, :, :, start_up:end_up - cr] += preds[mm][:, :, :, :-cr]\n            mask[:, :, :, start_up:end_up - cr] += 1\n        elif mm == (len(starts) - 1):\n            out[:, :, :, start_up + cr:end_up] += preds[mm][:, :, :, cr:]\n            mask[:, :, :, start_up + cr:end_up] += 1\n        else:\n            out[:, :, :, start_up + cr:end_up - cr] += preds[mm][:, :, :, cr:-cr]\n            mask[:, :, :, start_up + cr:end_up - cr] += 1\n\n    out /= mask\n    return out\ndef sliding_window_half_segmentation(\n    image,\n    models,\n    window_size: int = 528*2,\n    overlap_frac: float = 0.25,\n    overlap_px: int = 256*2,\n    interp_mode: str = \"bilinear\",\n    align_corners: bool = True,\n    device = 'cuda',\n    scalesize = 2,\n    crop_remain = 64\n    \n):\n\n    image = image.unsqueeze(0)\n    B, C, H, W = image.shape\n    W = int(W / 2)\n    # Compute stride from overlap\n    if overlap_px is not None:\n        stride = window_size - overlap_px\n    else:\n        stride = max(1, int(round(window_size * (1 - overlap_frac))))\n    stride = max(1, stride)\n\n    # Compute window start indices covering the width\n    starts = []\n    ends = []\n    s = 0\n    while True:\n        starts.append(s)\n        ends.append(s + window_size)\n        if s + window_size >= W:\n            break\n        s_next = s + stride\n        # Ensure last window touches the end\n        if s_next + window_size > W:\n            s = W - window_size\n        else:\n            s = s_next\n        if starts and s == starts[-1]:\n            break\n\n    #image_up = F.interpolate(\n    #    image, size=(H, W * scalesize), mode=interp_mode, align_corners=align_corners\n    #)\n\n    preds = []\n    with torch.no_grad():\n        for s in starts:\n            \n            s_up = s * scalesize\n            e_up = (s + window_size) * scalesize\n            slice_img = image[:, :, :, s_up:e_up]\n            \n            inputs = {\"image\": slice_img}\n            pmy = None\n            for model in models:\n                pred = model(inputs)[\"pixel\"]\n                pmy = pred if pmy is None else (pmy + pred)\n            pmy /= len(models)\n            preds.append(pmy)\n\n    K = preds[0].shape[1]\n    out = torch.zeros((1, K, H, W * scalesize), dtype=preds[0].dtype, device=device)\n    mask = torch.zeros((1, 1, H, W * scalesize), dtype=preds[0].dtype, device=device)\n\n    for mm, (start, end) in enumerate(zip(starts, ends)):\n        start_up = start * scalesize\n        end_up = end * scalesize\n        cr = crop_remain * scalesize\n\n        if mm == 0:\n            out[:, :, :, start_up:end_up - cr] += preds[mm][:, :, :, :-cr]\n            mask[:, :, :, start_up:end_up - cr] += 1\n        elif mm == (len(starts) - 1):\n            out[:, :, :, start_up + cr:end_up] += preds[mm][:, :, :, cr:]\n            mask[:, :, :, start_up + cr:end_up] += 1\n        else:\n            out[:, :, :, start_up + cr:end_up - cr] += preds[mm][:, :, :, cr:-cr]\n            mask[:, :, :, start_up + cr:end_up - cr] += 1\n\n    out /= mask\n    return out\ndef full_segmentation(\n    image,\n    models,\n    interp_mode: str = \"bilinear\",\n    align_corners: bool = True,\n    device = 'cuda',\n    scalesize = 2,\n    \n):\n\n    image = image.unsqueeze(0)\n    B, C, H, W = image.shape\n    \n\n    #image_up = F.interpolate(\n    #    image, size=(H, W * scalesize), mode=interp_mode, align_corners=align_corners\n    #)\n\n    with torch.no_grad():\n        inputs = {\"image\": image}\n        pmy = None\n        for model in models:\n            pred = model(inputs)[\"pixel\"]\n            pmy = pred if pmy is None else (pmy + pred)\n        pmy /= len(models)\n \n    return pmy\n\ndef afterprocess_once(outputs,zero_mv,length,t0,t1,scale_size,mv_to_pixel,fill_model,predict_lenth,fs,th,al):\n    t0 = t0 * scale_size\n    t1 = t1 * scale_size\n    pixel = outputs.float().data.cpu().numpy()[0]\n    series_in_pixel,miss_in_pixel = stage2_common.pixel_to_series(pixel[..., t0:t1], zero_mv, length,th = th,tau = 0.0)\n    series_in_pixel,miss_in_pixel_ = fill_miss_runs(series_in_pixel, miss_in_pixel,zero_mv, min_len=5)\n    series_in_pixel =  np.array([\n                        ((zero_mv[i] - series_in_pixel[i]) / mv_to_pixel[i]) for i in range(4)\n                        ])\n    \n    #==================================\n    \"\"\"\n    if scale_size != 4:\n        series_in_pixel = torch.from_numpy(series_in_pixel).unsqueeze(1)\n        series_in_pixel = F.interpolate(series_in_pixel, size=predict_lenth*4, mode='linear', align_corners=al)\n        miss_in_pixel = torch.from_numpy(miss_in_pixel.astype(np.float32)).unsqueeze(1)\n        miss_in_pixel = F.interpolate(miss_in_pixel, size=predict_lenth*4, mode='nearest')\n        series_in_pixel = series_in_pixel.squeeze(1).cpu().numpy()\n        miss_in_pixel = miss_in_pixel.squeeze(1).cpu().numpy()\n    peaks = detect_global_rpeaks(series_in_pixel,fs = fs)\n    \n    series_in_pixel = fill_missing_by_phase_within_each_1250(\n                                            series_in_pixel, miss_in_pixel, peaks,\n                                            \n                                            total_len=series_in_pixel.shape[1],\n                                            seg_len=predict_lenth,\n                                           \n                                        )\n    \n    for i in range(4):\n        plt.plot(series_in_pixel[1][i*1969:(i+1)*1969])\n        for item in peaks:\n            if item <= (i+1)*1969:\n                if item > (i*1969):\n                    plt.axvline(x=item - 1969*i, color='r', linestyle='--')\n            else:\n                break\n        plt.show()\n    \n\n    series_in_pixel = torch.from_numpy(series_in_pixel).unsqueeze(0)\n    #series_in_pixel = F.interpolate(series_in_pixel, size=1969*2, mode='linear', align_corners=True)\n    miss_in_pixel = torch.from_numpy(miss_in_pixel.astype(np.float32)).unsqueeze(0)\n    #miss_in_pixel = F.interpolate(miss_in_pixel, size=1969*2, mode='nearest')\n    #series_in_pixel = series_in_pixel.squeeze(1).unsqueeze(0)\n    #miss_in_pixel = miss_in_pixel.squeeze(1).unsqueeze(0)\n    #miss_in_pixel[:,0,606*2:612*2] = 1\n    #miss_in_pixel[:,1,606*2:612*2] = 1\n    #miss_in_pixel[:,2,606*2:612*2] = 1\n    #miss_in_pixel[:,0,1099*2:1103*2] = 1\n    #miss_in_pixel[:,1,1099*2:1103*2] = 1\n    #miss_in_pixel[:,2,1099*2:1103*2] = 1\n    #miss_in_pixel[:,0,1591*2:1595*2] = 1\n    #miss_in_pixel[:,1,1591*2:1595*2] = 1\n    #miss_in_pixel[:,2,1591*2:1595*2] = 1\n    #series_in_pixel = series_in_pixel * (1.0 - miss_in_pixel)\n    with torch.no_grad():\n        res = fill_model(series_in_pixel.float().to(DEVICE) ,miss_in_pixel.float().to(DEVICE))[0]\n        series_in_pixel = res.cpu().numpy()\n    #==================================\n    \"\"\"\n    if length is not None and length != series_in_pixel.shape[1]:\n        series_in_pixel = torch.from_numpy(series_in_pixel).unsqueeze(1)\n        series_in_pixel = F.interpolate(series_in_pixel, size=length, mode='linear', align_corners=al)\n        series_in_pixel = series_in_pixel.squeeze(1).cpu().numpy()\n    series = series_in_pixel\n    return series\n\n\n\n\nfrom stage2_model import Net as Stage2Net_34\nfrom stage2_model_arunodhayan import Net as Stage2Net_18d\nimport stage2_common\nfrom scipy.signal import savgol_filter, medfilt\nos.makedirs(f'{OUT_DIR}/digitalised', exist_ok=True)\n#os.makedirs(f'{OUT_DIR}/debug', exist_ok=True)\n\ndef run_stage2(sample_id,image_1,image_2,\n               stage2_half_models,\n               stage2_halfcrop_models,\n               stage2_full_models,\n              stage2_full2_models):\n\n    scores = 0\n\n    length = read_sampling_length(sample_id) #5120\n    # at rectified coord frame: H, W = 1700, 2200\n    image_1 = cv2.resize(image_1[0:1696,0:2176], (2176*2, 1696), interpolation=cv2.INTER_LINEAR)\n    image_2 = cv2.resize(image_2[0:1696,0:2176], (2176*2, 1696), interpolation=cv2.INTER_LINEAR)\n    \n    #-------------------------------------------------------------------------------------------\n    crop_halfcrop = image_1[192:1696,0:2112*2,:]#image_1[192:1696, 0:2112]\n    #crop_halfcrop = cv2.resize(crop_halfcrop, (2112*2, 1504), interpolation=cv2.INTER_LINEAR)\n    crop_halfcrop = torch.from_numpy(np.ascontiguousarray(crop_halfcrop.transpose(2, 0, 1)))\n    outputs_halfcrop = sliding_window_half_segmentation(crop_halfcrop,stage2_halfcrop_models,)\n    series_halfcrop = afterprocess_once(outputs_halfcrop,[ 515, 799, 1081, 1341],length,\n                                        118,2086,2,[77.5, 79, 78, 77.5],fill_model,1969,787,0.2,False)\n    #-------------------------------------------------------------------------------------------\n    full_image = image_1#image_1[0:1696, 0:2176]\n    #full_image = cv2.resize(full_image, (W, H), interpolation=cv2.INTER_LINEAR)\n    full_image = torch.from_numpy(np.ascontiguousarray(full_image.transpose(2, 0, 1)))\n    outputs_full = full_segmentation(full_image,stage2_full_models,)\n    series_full = afterprocess_once(outputs_full,[ 706, 990, 1273, 1532],length,\n                                    118,2087,2,[78, 78, 78, 78],fill_model,1969,787,0.2,False)\n    #-------------------------------------------------------------------------------------------\n    crop_half = image_1[0:1696, 0:2112*2]\n    #crop_half = cv2.resize(crop_half, (2112*2, 1696), interpolation=cv2.INTER_LINEAR)\n    crop_half = torch.from_numpy(np.ascontiguousarray(crop_half.transpose(2, 0, 1)))\n    outputs_half = sliding_window_half_segmentation(crop_half,stage2_half_models,)\n    series_half = afterprocess_once(outputs_half,[707, 991, 1273, 1533],length,\n                                    118,2086,2,[77.5, 79, 78, 77.5],fill_model,1969,787,0.2,False)\n    \n    #-------------------------------------------------------------------------------------------\n    full_image2 = image_2\n    full_image2 = torch.from_numpy(np.ascontiguousarray(full_image2.transpose(2, 0, 1)))\n    outputs_full2 = full_segmentation(full_image2,stage2_full2_models,)\n    series_full2 = afterprocess_once(outputs_full2,[703.5, 987.5, 1271.5, 1531.5],length,\n                                     118,2080,2,[78.,78.,78.,78.],fill_model,1963,785,0.1,False)\n    #-------------------------------------------------------------------------------------------\n    series = 0.34*series_halfcrop + 0.34*series_full + 0.26*series_half + 0.06*series_full2\n    \n    tmp_ = (series[1][:int(series.shape[1] // 4)] + series[3][:int(series.shape[1] // 4)]) / 2.\n    series[1][:int(series.shape[1] // 4)] = tmp_\n    series[3][:int(series.shape[1] // 4)] = tmp_\n    #===============\n    L1_ = series[0][:int(series.shape[1] // 4)]\n    L2_ = series[1][:int(series.shape[1] // 4)]\n    L3_ = series[2][:int(series.shape[1] // 4)]\n    error_ = L2_ - (L1_ + L3_)\n    series[0][:int(series.shape[1] // 4)] = L1_ + (0.3 * error_)\n    series[1][:int(series.shape[1] // 4)] = L2_ - (0.3 * error_)\n    series[3][:int(series.shape[1] // 4)] = L2_ - (0.3 * error_)\n    series[2][:int(series.shape[1] // 4)] = L3_ + (0.3 * error_)\n    #===============\n    torch.cuda.empty_cache()\n    return series","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:55:16.20197Z","iopub.execute_input":"2026-01-20T07:55:16.202455Z","iopub.status.idle":"2026-01-20T07:55:16.285159Z","shell.execute_reply.started":"2026-01-20T07:55:16.202432Z","shell.execute_reply":"2026-01-20T07:55:16.284364Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#stage 0 models\nstage0_net = Stage0Net(pretrained=False)\nstage0_net = stage0_common.load_net(stage0_net, f'{WEIGHT_DIR}/stage0-last.checkpoint.pth')\nstage0_net.to(DEVICE)\n#stage1 models\nstage1_net = Stage1Net(pretrained=False)\nckp_stage1 = torch.load(f'/kaggle/input/hengck-libs/weight/stage1-last.checkpoint.pth',weights_only = False)['state_dict']\nstage1_net.load_state_dict(ckp_stage1)\nstage1_net.eval()\nstage1_net.output_type = ['infer']\nstage1_net.to(DEVICE)\n\n#stage2 models\n#resnet34\n#resnet34\nstage2_net_0116fold0 = Stage2Net_34(pretrained=False)\nckp_0116fold0 = torch.load(f'/kaggle/input/22jan-halfcrop/stage2-20260117_fold0_best_score.pth',weights_only = False)['state_dict']\nstage2_net_0116fold0.load_state_dict(ckp_0116fold0)\nstage2_net_0116fold0.eval()\nstage2_net_0116fold0.output_type = ['infer']\nstage2_net_0116fold0.to(DEVICE)\n\nstage2_net_0116fold1 = Stage2Net_34(pretrained=False)\nckp_0116fold1 = torch.load(f'/kaggle/input/22jan-halfcrop/stage2-20260117_fold1_best_score.pth',weights_only = False)['state_dict']\nstage2_net_0116fold1.load_state_dict(ckp_0116fold1)\nstage2_net_0116fold1.eval()\nstage2_net_0116fold1.output_type = ['infer']\nstage2_net_0116fold1.to(DEVICE)\n\n\n\nstage2_net_0116fold3 = Stage2Net_34(pretrained=False)\nckp_0116fold3 = torch.load(f'/kaggle/input/22jan-halfcrop/stage2-20260117_fold2_best_score.pth',weights_only = False)['state_dict']\nstage2_net_0116fold3.load_state_dict(ckp_0116fold3)\nstage2_net_0116fold3.eval()\nstage2_net_0116fold3.output_type = ['infer']\nstage2_net_0116fold3.to(DEVICE)\n\nstage2_net_0116fold4 = Stage2Net_34(pretrained=False)\nckp_0116fold4 = torch.load(f'/kaggle/input/ecg-waveseg-localweight/stage2-20260116_fold4_best_score.pth',weights_only = False)['state_dict']\nstage2_net_0116fold4.load_state_dict(ckp_0116fold4)\nstage2_net_0116fold4.eval()\nstage2_net_0116fold4.output_type = ['infer']\nstage2_net_0116fold4.to(DEVICE)\n#----------------\nstage2_net_0111fold0 = Stage2Net_34(pretrained=False)\nckp_0111fold0 = torch.load(f'/kaggle/input/ecg-waveseg-localweight/stage2-20260111_fold0_best_score.pth',weights_only = False)['state_dict']\nstage2_net_0111fold0.load_state_dict(ckp_0111fold0)\nstage2_net_0111fold0.eval()\nstage2_net_0111fold0.output_type = ['infer']\nstage2_net_0111fold0.to(DEVICE)\n\nstage2_net_0111fold1 = Stage2Net_34(pretrained=False)\nckp_0111fold1 = torch.load(f'/kaggle/input/ecg-waveseg-localweight/stage2-20260111_fold1_best_score.pth',weights_only = False)['state_dict']\nstage2_net_0111fold1.load_state_dict(ckp_0111fold1)\nstage2_net_0111fold1.eval()\nstage2_net_0111fold1.output_type = ['infer']\nstage2_net_0111fold1.to(DEVICE)\n\nstage2_net_0111fold2 = Stage2Net_34(pretrained=False)\nckp_0111fold2 = torch.load(f'/kaggle/input/ecg-waveseg-localweight/stage2-20260111_fold2_best_score.pth',weights_only = False)['state_dict']\nstage2_net_0111fold2.load_state_dict(ckp_0111fold2)\nstage2_net_0111fold2.eval()\nstage2_net_0111fold2.output_type = ['infer']\nstage2_net_0111fold2.to(DEVICE)\n\nstage2_net_0111fold3 = Stage2Net_34(pretrained=False)\nckp_0111fold3 = torch.load(f'/kaggle/input/ecg-waveseg-localweight/stage2-20260111_fold3_best_score.pth',weights_only = False)['state_dict']\nstage2_net_0111fold3.load_state_dict(ckp_0111fold3)\nstage2_net_0111fold3.eval()\nstage2_net_0111fold3.output_type = ['infer']\nstage2_net_0111fold3.to(DEVICE)\n\nstage2_net_0111fold4 = Stage2Net_34(pretrained=False)\nckp_0111fold4 = torch.load(f'/kaggle/input/ecg-waveseg-localweight/stage2-20260111_fold4_best_score.pth',weights_only = False)['state_dict']\nstage2_net_0111fold4.load_state_dict(ckp_0111fold4)\nstage2_net_0111fold4.eval()\nstage2_net_0111fold4.output_type = ['infer']\nstage2_net_0111fold4.to(DEVICE)\n\n#----------------resnet18d\nstage2_net_1225fold0_per = Stage2Net_18d(pretrained=False)\nckp_1225fold0 = torch.load(f'/kaggle/input/res18d-perlinenoise-allaug/300epochs_stage2-20251225_fold0_best_score.pth',weights_only = False)['state_dict']\nstage2_net_1225fold0_per.load_state_dict(ckp_1225fold0)\nstage2_net_1225fold0_per.eval()\nstage2_net_1225fold0_per.output_type = ['infer']\nstage2_net_1225fold0_per.to(DEVICE)\n\nstage2_net_1225fold1_per= Stage2Net_18d(pretrained=False)\nckp_1225fold1 = torch.load(f'/kaggle/input/res18d-perlinenoise-allaug/300epochs_stage2-20251225_fold1_best_score.pth',weights_only = False)['state_dict']\nstage2_net_1225fold1_per.load_state_dict(ckp_1225fold1)\nstage2_net_1225fold1_per.eval()\nstage2_net_1225fold1_per.output_type = ['infer']\nstage2_net_1225fold1_per.to(DEVICE)\n\nstage2_net_1225fold2_per = Stage2Net_18d(pretrained=False)\nckp_1225fold2 = torch.load(f'/kaggle/input/res18d-perlinenoise-allaug/300epochs_stage2-20251225_fold2_best_score.pth',weights_only = False)['state_dict']\nstage2_net_1225fold2_per.load_state_dict(ckp_1225fold2)\nstage2_net_1225fold2_per.eval()\nstage2_net_1225fold2_per.output_type = ['infer']\nstage2_net_1225fold2_per.to(DEVICE)\n\nstage2_net_1225fold3_per = Stage2Net_18d(pretrained=False)\nckp_1225fold3 = torch.load(f'/kaggle/input/res18d-perlinenoise-allaug/300epochs_stage2-20251225_fold3_best_score.pth',weights_only = False)['state_dict']\nstage2_net_1225fold3_per.load_state_dict(ckp_1225fold3)\nstage2_net_1225fold3_per.eval()\nstage2_net_1225fold3_per.output_type = ['infer']\nstage2_net_1225fold3_per.to(DEVICE)\n\nstage2_net_1225fold4_per = Stage2Net_18d(pretrained=False)\nckp_1225fold4 = torch.load(f'/kaggle/input/res18d-perlinenoise-allaug/300epochs_stage2-20251225_fold4_best_score.pth',weights_only = False)['state_dict']\nstage2_net_1225fold4_per.load_state_dict(ckp_1225fold4)\nstage2_net_1225fold4_per.eval()\nstage2_net_1225fold4_per.output_type = ['infer']\nstage2_net_1225fold4_per.to(DEVICE)\n\n\n\nstage2_net_fullori = Stage2Net_18d(pretrained=False)\nckp_fullori = torch.load(f'/kaggle/input/2xresolution-nomorphology-nocompression/pytorch/default/12/stage2_epoch300_0.00417.pth',weights_only = False)\nstage2_net_fullori.load_state_dict(ckp_fullori)\nstage2_net_fullori.eval()\nstage2_net_fullori.output_type = ['infer']\nstage2_net_fullori.to(DEVICE)\n#-----------------------------\n\n#-----------------------------\n\nstage2_halfcrop_models = [stage2_net_0111fold0,stage2_net_0111fold1,stage2_net_0111fold2,stage2_net_0111fold3,stage2_net_0111fold4]\nstage2_half_models = [stage2_net_0116fold0,stage2_net_0116fold1,stage2_net_0116fold3,stage2_net_0116fold4]\nstage2_full_models = [\n                     stage2_net_1225fold0_per,stage2_net_1225fold1_per,stage2_net_1225fold2_per,stage2_net_1225fold3_per,stage2_net_1225fold4_per]\nstage2_full2_models = [stage2_net_fullori]\n#postprocess models\nfill_model = ImputeTCNPeriodicMS(\n                                in_ch=4,\n                                hidden=128, k=7, p_drop=0.5,\n                                lstm_hidden=128, lstm_layers=1, lstm_dropout=0.5,\n                                k_harmonics=6,\n                                ms_branches=((1,2,4), (1,3,9), (2,4,8,16))\n                            )\nckp_fill_model = torch.load(f'/kaggle/input/ecg-wavefill-localweight/260105fold0_best_score.pth',weights_only = False)['model']\nfill_model.load_state_dict(ckp_fill_model)\nfill_model.eval()\nfill_model.to(DEVICE)\n\n\nfor sample_id in tqdm(valid_id,total = len(valid_id)):\n    \n    try:\n        normalised = run_stage0(stage0_net,sample_id)\n        rectified_image_1,rectified_image_2 = run_stage1(stage1_net,normalised)\n        series = run_stage2(sample_id,rectified_image_1,rectified_image_2\n                            ,stage2_half_models,stage2_halfcrop_models,stage2_full_models,stage2_full2_models)\n        np.save(f'{OUT_DIR}/digitalised/{sample_id}.series.npy', series)\n    except:\n        series = np.zeros((4,read_sampling_length(sample_id)))\n        np.save(f'{OUT_DIR}/digitalised/{sample_id}.series.npy', series)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:55:16.286144Z","iopub.execute_input":"2026-01-20T07:55:16.286376Z","iopub.status.idle":"2026-01-20T08:02:18.528639Z","shell.execute_reply.started":"2026-01-20T07:55:16.286357Z","shell.execute_reply":"2026-01-20T08:02:18.527787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODE == 'local':\n    scores = 0\n    for sample_id in tqdm(valid_id,total = len(valid_id)):\n        series = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.npy')\n        truth_df = stage2_common.read_truth_series(sample_id,KAGGLE_DIR)\n        truth_series = truth_df[['series0','series1','series2','series3',]].values.T\n        \"\"\"\n        t = np.arange(len(series[0]))\n        fig, axes = plt.subplots(4, 1, figsize=(12, 10))\n        for j in range(4):\n            snr=0\n            axes[j].plot(t, series[j], alpha=1.0, color='blue', linewidth=1, label='predict')\n            if MODE=='local':\n                axes[j].plot(t, truth_series[j], alpha=0.5, color='red', linewidth=1,label='truth')\n                snr = -stage2_common.np_snr(series[j], truth_series[j])\n        \n            axes[j].set_title(f'snr {snr:8.3f}')\n            axes[j].legend()\n        plt.show()\n        \"\"\"\n        for j in range(4):\n            snr=0\n            snr = -stage2_common.np_snr(series[j], truth_series[j])\n            scores += snr#print(snr)\n    scores /= 80\n    print(scores)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:02:18.529424Z","iopub.execute_input":"2026-01-20T08:02:18.529671Z","iopub.status.idle":"2026-01-20T08:02:18.770296Z","shell.execute_reply.started":"2026-01-20T08:02:18.529646Z","shell.execute_reply":"2026-01-20T08:02:18.769449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#make sbmission csv\n#FAIL_ID=[1053922973, ]\ndef series_dict(series):\n    d = {}\n    for l in range(3):\n        lead_names = [\n            ['I',   'aVR', 'V1', 'V4'],\n            ['II',  'aVL', 'V2', 'V5'],\n            ['III', 'aVF', 'V3', 'V6'],\n        ][l]\n        split = np.array_split(series[l], 4)\n        for (k, s) in zip(lead_names, split):\n            d[k] = s\n    \n    d['II'] = series[3]\n    \n    return d\ndef make_submission():\n\tprint('===========================================')\n\tprint('making submission csv ...')\n\n\tsubmit_df=[]\n\tgb = valid_df.groupby('id')\n\tfor i,(sample_id, df) in enumerate(gb):\n        \n\t\ttry:\n\t\t\tseries = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.npy')\n            \n\t\t\tseries_by_lead={}\n\t\t\tfor l in range(3):\n\t\t\t\tlead = [\n\t\t\t\t\t['I',   'aVR', 'V1', 'V4'],\n\t\t\t\t\t['II',  'aVL', 'V2', 'V5'],\n\t\t\t\t\t['III', 'aVF', 'V3', 'V6'],\n\t\t\t\t][l]\n\n\t\t\t\tlength=[\n\t\t\t\t\tdf[df['lead']==lead[j]].iloc[0].number_of_rows\n\t\t\t\t\tfor j in range(4)\n\t\t\t\t]\n\t\t\t\tif lead[0]=='II':\n\t\t\t\t\tlength[0] = length[0]-sum(length[1:])\n\n\t\t\t\tindex = np.cumsum(length)[:-1]\n\t\t\t\tsplit = np.split(series[l], index)\n\t\t\t\t#print(length)\n\t\t\t\tfor (k, s) in zip(lead, split):\n\t\t\t\t\tseries_by_lead[k] = s\n\t\t\t\t\t#print(k,len(s))\n\t\t\tseries_by_lead['II'] = series[3]\n\t\t\t#print(series_by_lead)\n   \n\t\texcept: \n\t\t\tseries_by_lead = {}\n\t\t\tfor j,d in df.iterrows():\n\t\t\t\tseries_by_lead[d.lead] = np.zeros(d.number_of_rows)\n\n\t\t#print('\\r\\t {sample_id}', end='', flush=True)\n\t\tfor j,d in df.iterrows():\n\t\t\t#print(d.lead, len(series_by_lead[d.lead]),d.number_of_rows)\n            \n\t\t\t#assert(len(series_by_lead[d.lead])==d.number_of_rows)\n            \n            \n            #probably error here ... ???\n\t\t\tseries_by_lead[d.lead] = np.concatenate([\n                series_by_lead[d.lead], np.zeros_like(series_by_lead[d.lead])\n            ])[:d.number_of_rows]\n\t\t\tassert(len(series_by_lead[d.lead])==d.number_of_rows) \n\t\t\tprint(f'\\r\\t {i} {sample_id} : {d.lead}', end='', flush=True)\n\n\t\t\trow_id = [\n\t\t\t\tf'{sample_id}_{i}_{d.lead}' for i in range(d.number_of_rows)\n\t\t\t]\n\t\t\tthis_df = pd.DataFrame({\n\t\t\t\t'id':row_id,\n\t\t\t\t'value': series_by_lead[d.lead].astype(np.float32),\n\t\t\t})\n\t\t\tsubmit_df.append(this_df)\n\n\tprint('')\n\tsubmit_df = pd.concat(submit_df, axis=0, ignore_index=True, sort=False, copy=False)\n\tprint(submit_df)\n\tsubmit_df.to_csv('submission.csv',index=False)\n\nif (MODE=='fake')|(MODE=='submit'):\n    make_submission()\n    print('make_submission() ok!!!\\n')\n    if MODE=='submit':\n        shutil.rmtree(OUT_DIR)\n    !ls\n    #!rm -rf {OUT_DIR}\n\n'''\nfake:\n[21618231 rows x 2 columns]\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:02:18.771255Z","iopub.execute_input":"2026-01-20T08:02:18.771556Z","iopub.status.idle":"2026-01-20T08:02:18.789177Z","shell.execute_reply.started":"2026-01-20T08:02:18.771532Z","shell.execute_reply":"2026-01-20T08:02:18.788517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}