{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":13746387,"sourceType":"datasetVersion","datasetId":8747012},{"sourceId":13816899,"sourceType":"datasetVersion","datasetId":8620533},{"sourceId":271051632,"sourceType":"kernelVersion"},{"sourceId":677607,"sourceType":"modelInstanceVersion","modelInstanceId":513841,"modelId":528480}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip uninstall -y tensorflow\n!uv pip install --no-deps --system --no-index --find-links='/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/setup' 'connected-components-3d'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-31T03:18:18.065357Z","iopub.execute_input":"2025-12-31T03:18:18.066136Z","iopub.status.idle":"2025-12-31T03:18:19.750321Z","shell.execute_reply.started":"2025-12-31T03:18:18.066101Z","shell.execute_reply":"2025-12-31T03:18:19.749346Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile constant.py\n\nimport warnings\nwarnings.filterwarnings(\"ignore\", category=UserWarning, module=\"pydantic\")\n\nimport kagglehub\nseed = 0\nCUDA0 = \"cuda:0\"\ndeterministic = kagglehub.package_import('wasupandceacar/deterministic').deterministic\ndeterministic.init_all(seed, disable_list=['cuda_block'])\n\nimport sys\nsys.path.append('/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet')\n\nimport os\nimport traceback\nfrom pathlib import Path\nfrom shutil import copyfile\nimport torch\nimport cv2\nimport pandas as pd\nimport numpy as np\nfrom tqdm.auto import tqdm\n\nif_submit = os.getenv('KAGGLE_IS_COMPETITION_RERUN')\n\nif if_submit:\n    test_meta = Path(\"/kaggle/input/physionet-ecg-image-digitization/test.csv\")\n    test_dir = Path(\"/kaggle/input/physionet-ecg-image-digitization/test\")\nelse:\n    test_meta = Path(\"/kaggle/input/physio-test-fake-dataset/test_fake/test.csv\")\n    test_dir = Path(\"/kaggle/input/physio-test-fake-dataset/test_fake\")\n\nvalid_df = pd.read_csv(test_meta)\nvalid_df['id'] = valid_df['id'].astype(str) \nvalid_id = valid_df['id'].unique().tolist()\n\nFLOAT_TYPE = torch.float32\n\nglobal_dict = {\n    \"stage0_dir\": \"/kaggle/working/stage0\",\n    \"stage1_dir\": \"/kaggle/working/stage1\",\n    \"stage2_dir\": \"/kaggle/working/stage2\",\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T03:18:19.752069Z","iopub.execute_input":"2025-12-31T03:18:19.752320Z","iopub.status.idle":"2025-12-31T03:18:19.758672Z","shell.execute_reply.started":"2025-12-31T03:18:19.752296Z","shell.execute_reply":"2025-12-31T03:18:19.758014Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile stage0.py\nfrom constant import *\nfrom stage0_model import Net as Stage0Net\nfrom stage0_common import *\n\nimport numpy as np\nimport cv2\n\ndef apply_grayscale_guidance(image_rgb: np.ndarray) -> np.ndarray:\n    \"\"\"\n    更稳健的灰度引导：\n    1) 灰度\n    2) 大核 median blur 估计背景，做 divide 实现光照归一化（对阴影/曝光不均更稳）\n    3) NLM 去噪\n    4) CLAHE\n    5) 转回 3 通道\n    \"\"\"\n    gray = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2GRAY)\n\n    h, w = gray.shape[:2]\n    # 自适应大核：跟图像尺寸相关，避免手工调参；保证 odd\n    k = max(31, (min(h, w) // 20) | 1)\n    bg = cv2.medianBlur(gray, k)\n    bg = np.clip(bg, 1, 255).astype(np.uint8)\n\n    # 光照归一化：抑制大范围阴影/渐变\n    norm = cv2.divide(gray, bg, scale=255)\n    norm = np.clip(norm, 0, 255).astype(np.uint8)\n\n    denoised = cv2.fastNlMeansDenoising(norm, h=10)\n    clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8, 8))\n    contrast = clahe.apply(denoised)\n\n    return cv2.cvtColor(contrast, cv2.COLOR_GRAY2RGB)\n\n\nstage0_dir = Path(global_dict[\"stage0_dir\"])\nstage0_dir.mkdir(exist_ok=True)\n\nstage0_net = Stage0Net(pretrained=False)\nstage0_net = load_net(\n    stage0_net,\n    '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/weight/stage0-last.checkpoint.pth'\n)\nstage0_net.to(CUDA0)\nstage0_net.eval()\n\nfor n, sample_id in enumerate(tqdm(valid_id)):\n    path = test_dir / f'{sample_id}.png'\n    output_path = stage0_dir / f'{sample_id}.png'\n\n    image_original = cv2.imread(str(path), cv2.IMREAD_COLOR)\n    if image_original is None:\n        # 读不到就跳过（也可以 copy 原图）\n        copyfile(path, output_path)\n        continue\n\n    image_original = cv2.cvtColor(image_original, cv2.COLOR_BGR2RGB)\n\n    # 用 guidance 喂模型，但几何操作仍作用在原图坐标系\n    image_for_model = apply_grayscale_guidance(image_original)\n    batch = image_to_batch(image_for_model)\n\n    try:\n        with torch.no_grad(), torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            output = stage0_net(batch)\n\n        rotated, keypoint = output_to_predict(image_original, batch, output)\n        normalised, _, _ = normalise_by_homography(rotated, keypoint)\n\n        cv2.imwrite(str(output_path), cv2.cvtColor(normalised, cv2.COLOR_RGB2BGR))\n    except Exception:\n        traceback.print_exc()\n        copyfile(path, output_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T03:18:19.759410Z","iopub.execute_input":"2025-12-31T03:18:19.759725Z","iopub.status.idle":"2025-12-31T03:18:19.781480Z","shell.execute_reply.started":"2025-12-31T03:18:19.759708Z","shell.execute_reply":"2025-12-31T03:18:19.780704Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile stage1.py\nfrom constant import *\n\nfrom stage1_model import Net as Stage1Net\nfrom stage1_common import *\n\nimport numpy as np\nimport cv2\n\ndef apply_grayscale_guidance(image_rgb: np.ndarray) -> np.ndarray:\n    gray = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2GRAY)\n\n    h, w = gray.shape[:2]\n    k = max(31, (min(h, w) // 20) | 1)\n    bg = cv2.medianBlur(gray, k)\n    bg = np.clip(bg, 1, 255).astype(np.uint8)\n\n    norm = cv2.divide(gray, bg, scale=255)\n    norm = np.clip(norm, 0, 255).astype(np.uint8)\n\n    denoised = cv2.fastNlMeansDenoising(norm, h=10)\n    clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8, 8))\n    contrast = clahe.apply(denoised)\n\n    return cv2.cvtColor(contrast, cv2.COLOR_GRAY2RGB)\n\n\nstage0_dir = Path(global_dict[\"stage0_dir\"])\nstage1_dir = Path(global_dict[\"stage1_dir\"])\nstage1_dir.mkdir(exist_ok=True)\n\nstage1_net = Stage1Net(pretrained=False)\nstage1_net = load_net(\n    stage1_net,\n    '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/weight/stage1-last.checkpoint.pth'\n)\nstage1_net.to(CUDA0)\nstage1_net.eval()\n\nfor n, sample_id in enumerate(tqdm(valid_id)):\n    path = stage0_dir / f'{sample_id}.png'\n    output_path = stage1_dir / f'{sample_id}.png'\n\n    # stage0 输出读入：这里明确读 BGR->RGB，避免 flag 兼容性问题\n    image_original = cv2.imread(str(path), cv2.IMREAD_COLOR)\n    if image_original is None:\n        copyfile(path, output_path)\n        continue\n    image_original = cv2.cvtColor(image_original, cv2.COLOR_BGR2RGB)\n\n    # 用 guidance 预测网格点，但 rectify 用原图（更接近 stage2 训练输入分布）\n    image_for_model = apply_grayscale_guidance(image_original)\n    batch = {\n        'image': torch.from_numpy(\n            np.ascontiguousarray(image_for_model.transpose(2, 0, 1))\n        ).unsqueeze(0)\n    }\n\n    try:\n        with torch.no_grad(), torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            output = stage1_net(batch)\n\n        gridpoint_xy, _ = output_to_predict(image_for_model, batch, output)\n        rectified = rectify_image(image_original, gridpoint_xy)\n\n        cv2.imwrite(str(output_path), cv2.cvtColor(rectified, cv2.COLOR_RGB2BGR))\n    except Exception:\n        traceback.print_exc()\n        copyfile(path, output_path)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T03:18:19.783319Z","iopub.execute_input":"2025-12-31T03:18:19.783895Z","iopub.status.idle":"2025-12-31T03:18:19.803285Z","shell.execute_reply.started":"2025-12-31T03:18:19.783878Z","shell.execute_reply":"2025-12-31T03:18:19.802550Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile stage2.py\nfrom constant import *\nimport torchvision.transforms as T\nfrom stage2_model import *\nfrom stage2_common import *\nimport traceback\nimport numpy as np\nimport cv2\nfrom pathlib import Path\nfrom tqdm import tqdm\nfrom scipy.signal import savgol_filter, medfilt\n\nPROB_SMOOTH_SIGMA_Y = 0.2   # 建议范围：0.0 ~ 1.5；设 0.0 等价于关闭\n\n# 2) 置信度阈值：低于该值认为该列不可靠，先做插值补洞\nCONF_THR = 0.08            # 建议范围：0.08 ~ 0.20\n\n# 3) 去尖刺（Hampel-like）\nDESPIKE_WINDOW = 7         # 建议范围：7 ~ 21（奇数）\nDESPIKE_N_SIGMAS = 2.5     # 建议范围：2.5 ~ 4.0\n\n\ndef _odd(k: int) -> int:\n    return k if (k % 2 == 1) else (k + 1)\n\ndef safe_savgol(x: np.ndarray, window_length=7, polyorder=2) -> np.ndarray:\n    n = len(x)\n    if n < polyorder + 3:\n        return x\n    wl = _odd(window_length)\n    if wl > n:\n        wl = n if (n % 2 == 1) else (n - 1)\n    if wl < polyorder + 3 or wl < 3:\n        return x\n    po = min(polyorder, wl - 1)\n    return savgol_filter(x, window_length=wl, polyorder=po, mode='interp')\n\ndef despike_hampel_fast(x: np.ndarray, window=9, n_sigmas=3.0) -> np.ndarray:\n    \"\"\"\n    向量化 Hampel-like：\n    median(x) + MAD 阈值检测尖刺，并用局部中值替换\n    \"\"\"\n    w = _odd(int(window))\n    if len(x) < w:\n        return x\n    med = medfilt(x, kernel_size=w)\n    diff = np.abs(x - med)\n    mad = medfilt(diff, kernel_size=w) * 1.4826 + 1e-6\n    out = diff > (n_sigmas * mad)\n    y = x.copy()\n    y[out] = med[out]\n    return y\n\ndef interp_by_conf(x: np.ndarray, conf: np.ndarray, thr: float) -> np.ndarray:\n    if conf is None:\n        return x\n    good = conf >= thr\n    if np.count_nonzero(good) < 2:\n        return x\n    idx = np.arange(len(x))\n    return np.interp(idx, idx[good], x[good])\n\ndef postprocess_series(series_mv: np.ndarray, conf: np.ndarray) -> np.ndarray:\n    \"\"\"\n    series_mv: (4, length)\n    conf:      (4, length) 每个采样点的置信度（越大越可靠）\n    \"\"\"\n    out = series_mv.astype(np.float32, copy=True)\n    for i in range(out.shape[0]):\n        x = out[i]\n        x = np.nan_to_num(x, nan=0.0, posinf=0.0, neginf=0.0)\n\n        # 1) 低置信度点先补洞（避免错误点进入平滑导致“拉扯”）\n        x = interp_by_conf(x, conf[i] if conf is not None else None, CONF_THR)\n\n        # 2) 去尖刺（对偶发跳变非常有效）\n        x = despike_hampel_fast(x, window=DESPIKE_WINDOW, n_sigmas=DESPIKE_N_SIGMAS)\n\n        # 3) 保留你原来的平滑逻辑（参数不变）\n        x = medfilt(x, kernel_size=3)\n        x = safe_savgol(x, window_length=7, polyorder=2)\n\n        out[i] = x\n    return out\n\n\nclass Net3(nn.Module):\n    def __init__(self, pretrained=True):\n        super(Net3, self).__init__()\n        encoder_dim = [64, 128, 256, 512]\n        decoder_dim = [128, 64, 32, 16]\n\n        self.encoder = timm.create_model(\n            model_name='resnet34.a3_in1k', pretrained=pretrained,\n            in_chans=3, num_classes=0, global_pool=''\n        )\n        self.decoder = MyCoordUnetDecoder(\n            in_channel=encoder_dim[-1],\n            skip_channel=encoder_dim[:-1][::-1] + [0],\n            out_channel=decoder_dim,\n            scale=[2, 2, 2, 2]\n        )\n        self.pixel = nn.Conv2d(decoder_dim[-1], 4, 1)\n\n    def forward(self, image):\n        encode = encode_with_resnet(self.encoder, image)\n        last, _ = self.decoder(feature=encode[-1], skip=encode[:-1][::-1] + [None])\n        return self.pixel(last)\n\n\n# 配置（保持不动）\nstage1_dir = Path(global_dict[\"stage1_dir\"])\nstage2_dir = Path(global_dict[\"stage2_dir\"])\nstage2_dir.mkdir(exist_ok=True)\n\nstage2_net = Net3(pretrained=False).to(CUDA0)\nmodel_path = \"/kaggle/input/physio-seg-public/pytorch/net3_009_4200/1/iter_0004200.pt\"\nstage2_net.load_state_dict(torch.load(model_path))\nstage2_net.eval()\n\nx0, x1 = 0, 2176\ny0, y1 = 0, 1696\nzero_mv = [703.5, 987.5, 1271.5, 1531.5]   # ✅ 按你要求：不改\nmv_to_pixel = 78.5                          # ✅ 不改\nt0, t1 = 235, 4161                          # ✅ 不改\n\nresize = T.Resize((1696, 4352), interpolation=T.InterpolationMode.BILINEAR)\n\n\nfor n, sample_id in enumerate(tqdm(valid_id)):\n    path = stage1_dir / f'{sample_id}.png'\n    output_path = stage2_dir / f'{sample_id}.npy'\n\n    try:\n        image = cv2.imread(str(path), cv2.IMREAD_COLOR)\n        if image is None:\n            raise RuntimeError(f\"failed to read {path}\")\n\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        # 防御性裁剪：避免偶发尺寸不一致引发越界\n        H, W = image.shape[:2]\n        xx0, xx1 = max(0, x0), min(W, x1)\n        yy0, yy1 = max(0, y0), min(H, y1)\n\n        length = valid_df[(valid_df['id'] == sample_id) & (valid_df['lead'] == 'II')].iloc[0].number_of_rows\n\n        crop_img = image[yy0:yy1, xx0:xx1].astype(np.float32) / 255.0\n        input_tensor = torch.from_numpy(np.ascontiguousarray(crop_img.transpose(2, 0, 1))).unsqueeze(0)\n        batch = resize(input_tensor).to(CUDA0)\n\n        with torch.no_grad(), torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            output = stage2_net(batch)\n\n        pixel = torch.sigmoid(output).detach().cpu().numpy()[0].astype(np.float32)  # (4,H,W)\n        if n == 0:\n            print(\"stage2 input image mean/std:\", crop_img.mean(), crop_img.std())\n            print(\"pixel prob stats: min/mean/max =\", pixel.min(), pixel.mean(), pixel.max())\n        # 可选：纵向平滑概率图（减少毛刺）\n        if PROB_SMOOTH_SIGMA_Y and PROB_SMOOTH_SIGMA_Y > 0:\n            for c in range(pixel.shape[0]):\n                pixel[c] = cv2.GaussianBlur(pixel[c], (1, 0),\n                                            sigmaX=0, sigmaY=float(PROB_SMOOTH_SIGMA_Y))\n\n        pixel_crop = pixel[..., t0:t1]  # (4,H,Wc)\n\n        # 置信度：每列 max_y(prob)，再插值到 length\n        conf_w = pixel_crop.max(axis=1)  # (4,Wc)\n        x_old = np.linspace(0, 1, conf_w.shape[1], dtype=np.float32)\n        x_new = np.linspace(0, 1, length, dtype=np.float32)\n        conf = np.stack([np.interp(x_new, x_old, conf_w[i]).astype(np.float32) for i in range(4)], axis=0)\n\n        # 原流程：pixel -> series_in_pixel -> mV\n        series_in_pixel = pixel_to_series(pixel_crop, zero_mv, length)\n        series_mv = (np.array(zero_mv).reshape(4, 1) - series_in_pixel) / mv_to_pixel\n\n        # 新增：更稳健后处理（补洞 + 去尖刺 + 原平滑）\n        series_mv = postprocess_series(series_mv, conf)\n\n        np.save(output_path, series_mv)\n\n    except Exception:\n        traceback.print_exc()\n        raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T03:18:19.804122Z","iopub.execute_input":"2025-12-31T03:18:19.804285Z","iopub.status.idle":"2025-12-31T03:18:19.827579Z","shell.execute_reply.started":"2025-12-31T03:18:19.804272Z","shell.execute_reply":"2025-12-31T03:18:19.826855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python stage0.py\n!python stage1.py\n!python stage2.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T03:18:19.828300Z","iopub.execute_input":"2025-12-31T03:18:19.829122Z","iopub.status.idle":"2025-12-31T03:20:23.783163Z","shell.execute_reply.started":"2025-12-31T03:18:19.829104Z","shell.execute_reply":"2025-12-31T03:20:23.782117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nfrom pathlib import Path\nimport random\n\nstage2_dir = Path(\"/kaggle/working/stage2\")\nfiles = list(stage2_dir.glob(\"*.npy\"))\n\nif len(files) > 0:\n    \n    sample_file = random.choice(files)\n    \n    print(f\"📊 Analyzing file: {sample_file.name}\")\n    \n    \n    series = np.load(sample_file)\n    \n    fig, axes = plt.subplots(4, 1, figsize=(18, 12), sharex=True)\n    lead_names = [\"Leads Group 1 (I, II, III)\", \"Leads Group 2 (aVR, aVL, aVF)\", \"Leads Group 3 (V1-V6)\", \"Long Lead II\"]\n    \n    for i in range(4):\n        ax = axes[i]\n        signal = series[i, :]\n        \n        ax.plot(signal, color='#1f77b4', linewidth=1.2)\n        ax.set_title(f\"{lead_names[i]} | Min: {signal.min():.2f} | Max: {signal.max():.2f}\", fontsize=12)\n        ax.grid(True, alpha=0.3)\n        ax.set_ylabel(\"Voltage (mV)\")\n        \n        ax.axhline(0, color='red', linestyle='--', alpha=0.5, linewidth=0.8)\n\n    plt.xlabel(\"Time (Samples)\")\n    plt.tight_layout()\n    plt.show()\n    \n    print(\"✅ Done plotting.\")\nelse:\n    print(\"❌ No output files found in stage2!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T03:20:23.784566Z","iopub.execute_input":"2025-12-31T03:20:23.785265Z","iopub.status.idle":"2025-12-31T03:20:24.579605Z","shell.execute_reply.started":"2025-12-31T03:20:23.785222Z","shell.execute_reply":"2025-12-31T03:20:24.578796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from constant import *\nimport gc\nfrom scipy.signal import resample\nimport pandas as pd\nimport numpy as np\nfrom tqdm import tqdm\nfrom pathlib import Path\nimport plotly.graph_objects as go\nfrom plotly.subplots import make_subplots\nfrom scipy.stats import pearsonr\nfrom sklearn.metrics import mean_squared_error\nfrom scipy.signal import correlate\n\ndef align_signals(pred, gt, fs, max_shift_sec=0.2):\n    max_shift = int(max_shift_sec * fs)\n    \n    lags = np.arange(-max_shift, max_shift + 1)\n    corr = correlate(pred - np.mean(pred), gt - np.mean(gt), mode='same')\n    # 找到中心点附近的峰值\n    center = len(corr) // 2\n    search_range = corr[center - max_shift : center + max_shift + 1]\n    best_lag = lags[np.argmax(search_range)]\n    # 1. 执行时间偏移\n    aligned_pred = np.roll(pred, best_lag)\n    # 处理边缘滚动带来的噪声（简单置为边缘值）\n    if best_lag > 0: aligned_pred[:best_lag] = aligned_pred[best_lag]\n    elif best_lag < 0: aligned_pred[best_lag:] = aligned_pred[best_lag-1]\n    # 2. 执行垂直对齐（移除均值差）\n    aligned_pred = aligned_pred - (np.mean(aligned_pred) - np.mean(gt))\n    return aligned_pred\n\ndef calculate_official_snr(pred_12, gt_12):\n    signal_power = np.sum(gt_12**2)\n    noise_power = np.sum((gt_12 - pred_12)**2)\n    # 防止除以 0\n    if noise_power == 0: return 100.0 \n    snr = 10 * np.log10(signal_power / noise_power)\n    return snr\n\ndef calculate_prd(pred, gt):\n    diff = np.sum((gt - pred)**2)\n    ref = np.sum(gt**2)\n    print(f\"{diff:<10} | {ref:<12.4f} | {pred:<12.4f} |  {gt:<12.4f} \")\n    return np.sqrt(diff / ref) * 100\n    \ndef show_pred_gt(pred, gt):\n    fig = make_subplots(rows=12, cols=1, subplot_titles=[f'Lead {i+1}' for i in range(12)])\n    for i in range(12):\n        fig.add_trace(go.Scatter(y=pred[:, i], mode='lines', name=f'Pred Lead {i+1}', line=dict(color='blue')), row=i+1, col=1)\n        fig.add_trace(go.Scatter(y=gt[:, i], mode='lines', name=f'GT Lead {i+1}', line=dict(color='red')), row=i+1, col=1)\n    fig.update_layout(height=1200, showlegend=False)\n    fig.show(renderer='iframe')\n    \ndef expand_4_to_12(pred4):\n    pred12 = np.zeros((pred4.shape[0], 12))\n    quat = pred4.shape[0] // 4\n\n    pred12[:quat, 0] = pred4[:quat, 0]\n    pred12[:, 1] = pred4[:, 3]\n    pred12[:quat, 2] = pred4[:quat, 2]\n    pred12[quat:2*quat, 3] = pred4[quat:2*quat, 0]\n    pred12[quat:2*quat, 4] = pred4[quat:2*quat, 1]\n    pred12[quat:2*quat, 5] = pred4[quat:2*quat, 2]\n    pred12[2*quat:3*quat, 6] = pred4[2*quat:3*quat, 0]\n    pred12[2*quat:3*quat, 7] = pred4[2*quat:3*quat, 1]\n    pred12[2*quat:3*quat, 8] = pred4[2*quat:3*quat, 2]\n    pred12[3*quat:4*quat, 9] = pred4[3*quat:4*quat, 0]\n    pred12[3*quat:4*quat, 10] = pred4[3*quat:4*quat, 1]\n    pred12[3*quat:4*quat, 11] = pred4[3*quat:4*quat, 2]\n    \n    return pred12\n\ndef series_dict(series):\n    series_by_lead = dict()\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            series_by_lead[k] = s\n    series_by_lead['II'] = series[3]\n    return series_by_lead\n    \ndef calculate_metrics(pred_12, gt_12):\n    \"\"\"\n    计算 12 导联的平均 RMSE 和相关系数\n    pred_12: (N, 12) 预测值\n    gt_12: (N, 12) 真值\n    \"\"\"\n    rmse_list = []\n    corr_list = []\n    # 逐导联计算\n    for i in range(12):\n        p = pred_12[:, i]\n        g = gt_12[:, i]\n        # 1. 计算 RMSE\n        rmse = np.sqrt(mean_squared_error(g, p))\n        rmse_list.append(rmse)\n        # 2. 计算相关系数 (处理全零信号的情况)\n        if np.std(p) == 0 or np.std(g) == 0:\n            corr = 0\n        else:\n            corr, _ = pearsonr(g, p)\n        corr_list.append(corr)\n        \n    return rmse_list, corr_list\n\nstage2_dir = Path(global_dict[\"stage2_dir\"])\n\nsubmit_df = list()\ngb = valid_df.groupby('id')\n\nshow = True\n\nprint(\"Generating submission file...\")\n\nfor rec_idx, (sample_id, df) in enumerate(tqdm(gb)):\n    try:\n        \n        series = np.load(stage2_dir / f'{sample_id}.npy')\n        if not if_submit :\n            all_rmse = []\n            all_corr = []\n            pred_raw = np.transpose(series, axes=(1, 0))\n            pred_12 = expand_4_to_12(pred_raw)\n            gt_path = \"/kaggle/input/physio-test-fake-dataset/7663343_inp1.csv\"\n            gt_df = pd.read_csv(gt_path).fillna(0)\n            gt_12 = gt_df.values\n\n            fs = 1000 \n            # 对 12 个导联逐一进行官方要求的对齐\n            aligned_pred_12 = np.zeros_like(pred_12)\n            for i in range(12):\n                aligned_pred_12[:, i] = align_signals(pred_12[:, i], gt_12[:, i], fs)\n            # 计算官方 SNR\n            final_snr = calculate_official_snr(aligned_pred_12, gt_12)\n            # 计算相关系数和 PRD (形态相似度)\n            corrs = [pearsonr(aligned_pred_12[:, i], gt_12[:, i])[0] for i in range(12)]\n            print(f\"--- Validation Result for {sample_id} ---\")\n            print(f\"Official SNR: {final_snr:.2f} dB\")\n            print(f\"Mean PRD: {np.mean(corrs):.4f}%\")\n                    # 计算指标\n            rmse_leads, corr_leads = calculate_metrics(pred_12, gt_12)\n\n            print(f\"{'Lead':<10} | {'RMSE (mV)':<12} | {'Correlation':<12}\")\n            print(\"-\" * 40)\n            leads_names = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n            for i, name in enumerate(leads_names):\n                print(f\"{name:<10} | {rmse_leads[i]:<12.4f} | {corr_leads[i]:<12.4f}\")\n            \n            print(\"-\" * 40)\n            print(f\"Overall Mean RMSE: {np.mean(rmse_leads):.4f} mV\")\n            print(f\"Overall Mean Correlation: {np.mean(corr_leads):.4f}\")\n            if show:\n                show_pred_gt(aligned_pred_12, gt_12)\n                show = False    \n        series_by_lead = series_dict(series)\n\n        for _, d in df.iterrows():\n            s = series_by_lead.get(d.lead, np.zeros(d.number_of_rows))\n            \n            if len(s) != d.number_of_rows:\n                x_old = np.linspace(0, 1, len(s))\n                x_new = np.linspace(0, 1, d.number_of_rows)\n                s = np.interp(x_new, x_old, s)\n            \n            row_id = [f'{sample_id}_{t}_{d.lead}' for t in range(d.number_of_rows)]\n            submit_df.append(pd.DataFrame({'id': row_id, 'value': s}))\n            \n    except Exception as e:\n        pass\n\n    if rec_idx % 100 == 0:\n        gc.collect()\n\nif submit_df:\n    final_df = pd.concat(submit_df, axis=0, ignore_index=True, sort=False, copy=False)\n    final_df.to_csv('submission.csv', index=False)\n    print(f\"✅ Done! Saved submission.csv with shape: {final_df.shape}\")\n    print(final_df.head())\nelse:\n    print(\"❌ Error: No predictions were generated!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T03:20:24.580734Z","iopub.execute_input":"2025-12-31T03:20:24.581060Z","iopub.status.idle":"2025-12-31T03:20:27.419356Z","shell.execute_reply.started":"2025-12-31T03:20:24.581041Z","shell.execute_reply":"2025-12-31T03:20:27.418651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = pd.read_csv('submission.csv')\nprint(len(sub))\nsub.head(30)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T03:20:27.420195Z","iopub.execute_input":"2025-12-31T03:20:27.420503Z","iopub.status.idle":"2025-12-31T03:20:27.635449Z","shell.execute_reply.started":"2025-12-31T03:20:27.420482Z","shell.execute_reply":"2025-12-31T03:20:27.634712Z"}},"outputs":[],"execution_count":null}]}