{"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":101849,"databundleVersionId":13093295,"sourceType":"competition"}],"dockerImageVersionId":31090,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import glob\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n# import dotenv\nimport torch\nimport pyarrow.parquet as pq\nimport torch.nn.functional as F\n\n# dotenv.load_dotenv()\nbase_url = \"/kaggle/input/ariel-data-challenge-2025\"\n\ndef show_figure(array, title=None, cmap=\"gray\", aspect=\"auto\"):\n    image = array\n    plt.figure(figsize=(8, 4))\n    plt.imshow(image, cmap=cmap, aspect=aspect, vmin=-840, vmax=-660)\n    if title:\n        plt.title(title)\n    plt.xlabel(\"Spectral direction — 356 px\")\n    plt.ylabel(\"Spatial direction — 32 px\")\n    plt.colorbar(label=\"Counts\")\n    plt.tight_layout()\n    plt.show()\n\n\ndef reverse_adc(raw_signal: torch.Tensor):\n    # NOTE: You will need to define AIRS_gain and AIRS_offset\n    # For this competition, they are the same for all planets:\n    AIRS_gain = 0.4369\n    AIRS_offset = -1000.0\n    return raw_signal / AIRS_gain + AIRS_offset  # Use division for reversal\n\n\ndef subtract_dark(raw_signal: torch.Tensor, dark_frame: torch.Tensor):\n    \"\"\"\n    Subtract the dark frame from the raw signal.\n    \"\"\"\n    # CORRECTED: Use a relative path to read axis_info\n    axis_info_path = os.path.join(base_url, \"axis_info.parquet\")\n    dt = pq.read_table(axis_info_path)[\"AIRS-CH0-integration_time\"].drop_null().to_numpy().copy()\n\n    # This logic for alternating integration times is specific to the competition\n    dt[1::2] += 0.1\n\n    dt = torch.tensor(dt, dtype=torch.float32, device=\"cuda\")\n    dt = dt.view(-1, 1, 1)\n    return raw_signal - dark_frame * dt\n\n\ndef fix_dead_pixels_vectorized(signal, dead_pixels):\n    \"\"\"\n    使用高效的卷积操作替换坏点。\n\n    Args:\n        signal (torch.Tensor): 输入信号，形状为 (B, H, W)，例如 (1250, 32, 356)。\n        dead_pixels (torch.Tensor): 坏点掩码，形状为 (H, W)，例如 (32, 356)。\n\n    Returns:\n        torch.Tensor: 修复后的信号，形状与输入相同。\n    \"\"\"\n    # 确保输入是 PyTorch Tensors\n    # if not isinstance(signal, torch.Tensor):\n    #     signal = torch.tensor(signal, dtype=torch.float32)\n    # if not isinstance(dead_pixels, torch.Tensor):\n    #     dead_pixels = torch.tensor(dead_pixels, dtype=torch.float32)\n\n    # --- 1. 准备卷积 ---\n    # a. 将2D掩码转换为布尔型，并适配3D信号的形状\n    # mask 形状: (32, 356) -> (1, 32, 356)，以便广播到 (1250, 32, 356)\n    mask = (dead_pixels == 1).unsqueeze(0)\n\n    # b. 创建一个3x3的求和卷积核\n    # conv2d需要4D输入: (out_channels, in_channels, kH, kW)\n    kernel = torch.ones((1, 1, 3, 3), device=signal.device, dtype=signal.dtype)\n\n    # c. 信号需要是4D: (B, C, H, W)。我们的信号是(B, H, W)，所以增加一个channel维度\n    # signal_4d 形状: (1250, 32, 356) -> (1250, 1, 32, 356)\n    signal_4d = signal.unsqueeze(1)\n\n    # # --- 2. 计算邻居像素的和 ---\n    # # 使用反射填充，与你原始代码的意图一致\n    # # padding=1确保卷积后尺寸不变\n    # sum_of_9_pixels = F.conv2d(signal_4d, kernel, padding='same', padding_mode='reflect')\n\n    # --- 2. 【核心修正】分开执行填充和卷积 ---\n\n    # 步骤 2a: 手动进行 'reflect' 填充\n    # 我们需要在最后两个维度（H和W）的上下左右各填充1个像素\n    # pad元组格式: (pad_left, pad_right, pad_top, pad_bottom)\n    padded_signal_4d = F.pad(signal_4d, (1, 1, 1, 1), mode='reflect')\n\n    # 步骤 2b: 在已填充的张量上执行卷积，此时 padding 设置为 0 或 'valid'\n    # 'valid' 意味着不进行任何填充，这正是我们现在需要的\n    sum_of_9_pixels = F.conv2d(padded_signal_4d, kernel, padding='valid')\n\n    # ----------------------------------------------\n\n    # 从4D转回3D，方便后续计算\n    sum_of_9_pixels = sum_of_9_pixels.squeeze(1)\n\n    # --- 3. 计算邻居像素的均值 ---\n    # 9个像素的和 - 中心像素 = 8个邻居的和\n    sum_of_8_neighbors = sum_of_9_pixels - signal\n    mean_of_8_neighbors = sum_of_8_neighbors / 8.0\n\n    # --- 4. 使用掩码替换坏点 ---\n    # torch.where是最高效、最清晰的条件替换方法\n    # torch.where(condition, value_if_true, value_if_false)\n    # mask会被自动广播到signal的形状\n    fixed_signal = torch.where(mask, mean_of_8_neighbors, signal)\n\n    return fixed_signal\n\n\ndef linearity_correction(signal, linear_corr):\n    \"\"\"\n    Perform linearity correction on the signal using polynomial coefficients and Horner's Method.\n    \"\"\"\n    # Reshape linear_corr to match the signal shape and use Horner's method\n    # Your implementation is good, but the one from the high-perf script is standard:\n    c = linear_corr.view(6, signal.shape[1], signal.shape[2])  # Ensure correct shape\n\n    # Using horner's method to evaluate the polynomial\n    return ((((c[5] * signal + c[4]) * signal + c[3]) * signal + c[2]) * signal + c[1]) * signal + c[0]\n\n\n# --- NEW: Add the CDS function ---\ndef get_cds(signal: torch.Tensor):\n    \"\"\"\n    Step 5: Get Correlated Double Sampling (CDS)\n    The science frames are alternating between the start of the exposure and the end of\n    the exposure. The final CDS is the difference (End of exposure) - (Start of exposure).\n    \"\"\"\n    if signal.ndim == 3:\n        return signal[1::2, :, :] - signal[::2, :, :]\n    else:\n        return signal[1::2,:] - signal[::2,:]\n\n\ndef flat_correction(signal, flat_frame):\n    \"\"\"\n    Perform flat field correction on the signal using the flat frame.\n    \"\"\"\n    # Normalize the flat frame\n    # A small epsilon is added to prevent division by zero\n    epsilon = 1e-8\n    normalized_flat = flat_frame / (torch.mean(flat_frame) + epsilon)\n\n    # Apply flat field correction\n    corrected_signal = signal / normalized_flat\n    return corrected_signal\n\n\n# --- REORDERED & UNCOMMENTED: The main processing pipeline ---\ndef process_data(raw_signal, calibration_data):\n    \"\"\"\n    Process the raw signal with the complete calibration pipeline.\n    \"\"\"\n    # Step 1: Reverse ADC\n    if raw_signal.ndim == 3:  # For AIRS-CH0\n        processed_signal = reverse_adc(raw_signal)\n\n        # Step 2: Substitute dead pixels\n        # Consider replacing with NaN masking for better results before interpolation\n        dead_pixels = calibration_data['dead']\n        processed_signal = fix_dead_pixels_vectorized(processed_signal, dead_pixels)\n\n        # Step 3: Perform linearity correction\n        linear_corr = calibration_data['linear_corr']\n        processed_signal = linearity_correction(processed_signal, linear_corr)\n\n        # Step 4: Subtract dark frame\n        dark_frame = calibration_data['dark']\n        processed_signal = subtract_dark(processed_signal, dark_frame)\n\n        # Step 5: Get Correlated Double Sampling (CDS) - THE NEW STEP\n        processed_signal = get_cds(processed_signal)\n\n        # Step 6: Perform flat field correction\n        flat_frame = calibration_data['flat']\n        processed_signal = flat_correction(processed_signal, flat_frame)\n\n        return processed_signal\n\n    else:\n        processed_signal = reverse_adc(raw_signal)\n        # processed_signal = fix_dead_pixels_vectorized(processed_signal, dead_pixels)\n        # processed_signal = linearity_correction(processed_signal, calibration_data['linear_corr'])\n        # processed_signal = subtract_dark(processed_signal, calibration_data['dark'])\n        # processed_signal = flat_correction(processed_signal, calibration_data['flat'])\n        processed_signal = get_cds(processed_signal)\n        return processed_signal\n\ndef read_signal(path):\n    \"\"\"\n    Read data from a parquet file and return a PyTorch tensor.\n    \"\"\"\n    # Read the parquet file using PyArrow\n    arrow_table = pq.read_table(path)\n\n    # Convert to Pandas DataFrame and then to NumPy array\n    numpy_array = arrow_table.to_pandas(zero_copy_only=False).to_numpy()\n\n    # Convert to PyTorch tensor\n    if \"AIRS-CH0\" in path:\n        tensor = torch.tensor(numpy_array, dtype=torch.float32, device=\"cuda\").reshape(-1, 32, 356)\n    else:\n        tensor = torch.tensor(numpy_array, dtype=torch.float32, device=\"cuda\")\n\n    return tensor\n\ndef read_calibration(calibration_base_path):\n    \"\"\"\n    Read calibration data from a parquet file and return a PyTorch tensor.\n    \"\"\"\n    calibration_data=dict()\n    calibration_name = [\"dark\", \"dead\", \"linear_corr\", \"flat\"]\n    for name in calibration_name:\n        path = os.path.join(calibration_base_path, f\"{name}.parquet\")\n        # Read the parquet file using PyArrow\n        arrow_table = pq.read_table(path)\n\n        # Convert to Pandas DataFrame and then to NumPy array\n        numpy_array = arrow_table.to_pandas(zero_copy_only=False).to_numpy()\n\n        # Convert to PyTorch tensor\n        tensor = torch.tensor(numpy_array, dtype=torch.float32, device=\"cuda\")\n        calibration_data[name] = tensor\n\n    return calibration_data\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-08-26T18:29:15.766354Z","iopub.execute_input":"2025-08-26T18:29:15.76655Z","iopub.status.idle":"2025-08-26T18:29:20.891909Z","shell.execute_reply.started":"2025-08-26T18:29:15.766531Z","shell.execute_reply":"2025-08-26T18:29:20.891168Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Dataset**# # ","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset\n# import preprocess\n\nAIR_SLICE = (39,321)\n# WIN_FGS = (23500, 44000)\n# WIN_AIRS = (1958, 3666)\n# WIN_FGS_OOT = 200\n# WIN_AIRS_OOT = 170\nBIN_LEN = 187 #时间分箱长度（作者使用）\nCUT_BEGIN, CUT_END = 75, 115 # in-transit固定窗口 (作者使用）\n\ndef resample_1d(arr:torch.Tensor, target_len: int=187)->torch.Tensor:\n    \"\"\"\n        将 1D 序列按均匀分箱重采样到 target_len。\n        Parameters\n        ----------\n        arr : torch.Tensor, shape (T,)\n            原始时间序列（例如白光）\n        target_len : int\n            目标分箱长度（默认 187）\n        Returns\n        -------\n        torch.Tensor, shape (target_len,)\n            分箱平均后的 1D 序列\n        \"\"\"\n    T= arr.shape[0]\n    m = T//target_len\n    if m < 1:\n        raise ValueError(f\"Cannot resample {T} to {target_len} (m={m})\")\n    trimmed = arr[:m*target_len].view(target_len,m).mean(dim=1)\n    return trimmed\n\ndef resample_2d_time(mat:torch.Tensor, bins:int=BIN_LEN)->torch.Tensor:\n    \"\"\"\n        将 2D 矩阵按“时间维”均匀分箱到 bins。\n        Parameters\n        ----------\n        mat : torch.Tensor, shape (T, W)\n            例如 AIRS 的 (time, wavelength) 图\n        bins : int\n            目标时间分箱长度（默认 187）\n        Returns\n        -------\n        torch.Tensor, shape (bins, W)\n            分箱平均后的 2D 矩阵\n        \"\"\"\n    T,m = mat.shape[0],mat.shape[0]//bins\n    trim = mat[:m*bins].view(bins,m,-1).mean(dim=1)\n    return trim\n\n# Dataset class for CNN model\nclass ArielCNNDataset(Dataset):\n    \"\"\"\n    读取并处理单个行星的观测，构造 1D/2D CNN 需要的输入与标签。\n    输出字段（keys 不变）：\n    - pid   : str\n    - wc (white-light curve)   : torch.float32, shape (1, 187)\n              * 对齐原作者：仅用 AIRS 构造白光\n              * 步骤：AIRS 沿 y 求和 → (T,356) → 切 39..321 → (T,283)\n                      → time-binning 到 187 → (187,283)\n                      → white_curve = sum_over_lambda / mean_over_all_pixels\n              * 不在 Dataset 内做 per-sample min-max。训练阶段用“训练集全局 min/max”。\n    - map2d : torch.float32, shape (1, 40, 283)\n              * 对齐原作者：做星光谱归一化（头 50 + 尾 50 帧），\n                切 in-transit 固定窗口 [75:115) 40 帧，\n                对整块做“去均值”（减去一个标量），最后加通道维。\n              * [-1,1] 的归一化也放在训练阶段（用训练集统计量）。\n    - （仅 train 时）\n      y_mean  : torch.float32, shape (1,)   —— 光谱均值（1D-CNN 目标）\n      y_shift : torch.float32, shape (283,) —— 去均值后的光谱（2D-CNN 目标）\n    Notes\n    -----\n    * 快路径（disk_cache_dir）期望缓存中的 wc 为形状 (1,187)，map2d 为 (1,40,283)。\n      如果你还在用旧缓存（wc 是 (1,374)），会抛出明确的错误提示。\n    \"\"\"\n    def __init__(self,base_dir,planet_ids=None,split=\"train\",return_labels=True,cache_calib=True,disk_cache_dir=None):\n        self.base_dir = base_dir\n        self.split = split\n        self.return_labels = return_labels and (split == \"train\")\n        self.cache_calib = cache_calib\n        self._calib_cache = {}\n        self.disk_cache_dir = disk_cache_dir\n\n        # planet list\n        if planet_ids is None:\n            split_dir = os.path.join(base_dir,split)\n            planet_ids = sorted(d for d in os.listdir(split_dir) if os.path.isdir(os.path.join(split_dir,d)))\n        self.pids = planet_ids\n\n        if self.return_labels:\n            df = pd.read_csv(os.path.join(base_dir, \"train.csv\")).set_index(\"planet_id\")\n            df.index =df.index.map(str) #保证与目录名一致\n            self._y_map = {str(idx):row.values.astype(\"float32\") for idx, row in df.iterrows()}\n\n    # # ---- 低级 I/O：读 parquet + 标定 → 返回已做 CDS 的张量 ----\n    def _load_sensor(self, pid, band):\n        \"\"\"\n                Parameters\n                ----------\n                pid : str\n                band: str, in {\"AIRS-CH0\", \"FGS1\"}\n                Returns\n                -------\n                torch.Tensor\n                    AIRS : (5625, 32, 356)\n                    FGS1 : (67500, 1024)\n                \"\"\"\n        pdir = os.path.join(self.base_dir, self.split, pid)\n        sig_path = os.path.join(pdir, f\"{band}_signal_0.parquet\")\n        cal_dir = os.path.join(pdir, f\"{band}_calibration_0\")\n\n        raw = read_signal(sig_path)\n\n        if self.cache_calib and (pid, band) in self._calib_cache:\n            calib = self._calib_cache[(pid, band)]\n        else:\n            calib = read_calibration(cal_dir)\n            if self.cache_calib:\n                self._calib_cache[(pid, band)] = calib\n\n        processed_signal = process_data(raw, calib)# AIRS:(5625,32,356) / FGS1:(67500,1024)\n        return processed_signal\n\n    def __len__(self): return len(self.pids)\n\n    def __getitem__(self, idx):\n        pid = self.pids[idx]\n\n        # 加个快速通道：如果有缓存（from .npz cache)直接读并返回：\n        if self.disk_cache_dir is not None:\n            npz_path = os.path.join(self.disk_cache_dir, f\"{pid}.npz\")\n            if os.path.exists(npz_path):\n                z = np.load(npz_path)\n                sample = {\n                    \"pid\": pid,\n                    \"wc\": torch.tensor(z[\"wc\"], dtype=torch.float32),  # (1, 187)\n                    \"map2d\": torch.tensor(z[\"map2d\"], dtype=torch.float32),  # (1, 40, 283)\n                }\n\n                # 形状检查\n                if sample[\"wc\"].shape!= (1,187):\n                    raise ValueError(\n                        f\"[Cache shape mismatch] wc shape {tuple(sample['wc'].shape)} != (1,187). \"\n                        f\"请切换到新缓存目录（如 cnn_cache_wc187/train），或重建缓存。\"\n                    )\n                if sample[\"map2d\"].shape != (1, 40, 283):\n                    raise ValueError(\n                        f\"[Cache shape mismatch] map2d shape {tuple(sample['map2d'].shape)} != (1,40,283).\"\n                    )\n\n                if self.return_labels:\n                    sample.update({\n                        \"y_mean\": torch.tensor(z[\"y_mean\"], dtype=torch.float32),  # (1,)\n                        \"y_shift\": torch.tensor(z[\"y_shift\"], dtype=torch.float32)  # (283,)\n                    })\n                return sample\n        # 慢路径：现算\n        # -----FGS 白光：像素和-》（67500，）暂时用不着，先封起来\n        # fgs = self._load_sensor(pid, \"FGS1\") #(67500,1024)\n        # fgs_white = fgs.sum(dim=1).cpu()# (67500,)\n\n        # --- AIRS: 沿 y 求和 -> (T,356)，再切 39–321 -> (T,283)\n        airs = self._load_sensor(pid,\"AIRS-CH0\") #(5625,32,356)\n        airs_2d = airs.sum(dim=1)[:,AIR_SLICE[0]:AIR_SLICE[1]+1].cpu() # (5625,283)\n\n        # === 1D-CNN 的 wc：对齐原作者（仅 AIRS，(1,187)） ===\n        # time-binning 到 187 帧\n        airs_2d_bin = resample_2d_time(airs_2d) # (187,283)\n        # 按作者：white_curve = sum_over_lambda/mean_over_all_pixels\n        wc_mean = airs_2d_bin.mean()\n        wc_187 = airs_2d_bin.sum(dim=1)/ (wc_mean+1e-8) # (187,)\n        wc = wc_187.unsqueeze(0) # (1,187)，不做min-max 归一化per sample\n\n        # === 2D-CNN 的 map2d：星光谱归一化 → in-transit 切片 → 整块去均值 ===\n        # 星光谱：头 50 + 尾 50 帧（对齐作者 norm_star_spectrum）\n        oot_left = airs_2d_bin[:50].mean(dim=0) # (283,)\n        oot_right = airs_2d_bin[-50:].mean(dim=0) # (283,)\n        star_spec = oot_left + oot_right # (283,)\n        airs_2d_norm = airs_2d_bin/(star_spec.clamp_min(1e-8)) # (187,283)\n        # in-transit 固定窗口 [75:115) 40 帧\n        map_slice = airs_2d_norm[CUT_BEGIN:CUT_END] # (40,283)\n        # 整块去均值（减去一个标量）\n        map_centered = map_slice - map_slice.mean() # (40,283)\n        # 加通道维度，得到 (1,40,283)\n        map2d = map_centered.unsqueeze(0) # (1,40,283)\n\n        sample = {\"pid\":pid, \"wc\":wc, \"map2d\":map2d}\n\n        if self.return_labels:\n            y_full = torch.from_numpy(self._y_map[pid]) # (283,)\n            y_mean = y_full.mean().unsqueeze(0) # (1,)\n            sample.update({\n                \"y_mean\": y_mean,\n                \"y_shift\": y_full - y_mean,\n            })\n\n        return sample","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T18:29:20.893187Z","iopub.execute_input":"2025-08-26T18:29:20.893566Z","iopub.status.idle":"2025-08-26T18:29:20.994249Z","shell.execute_reply.started":"2025-08-26T18:29:20.893541Z","shell.execute_reply":"2025-08-26T18:29:20.993627Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **1D CNN**# ","metadata":{}},{"cell_type":"code","source":"\"\"\"\nOne-Dimensional CNN for Ariel ADC 2025 (PyTorch, aligned with starter Keras)\n============================================================================\n\n输入  : 归一化白光曲线 wc — 形状 **(batch, 1, 187)**\n         （仅 AIRS 构造：sum_λ / mean_all；不做 per-sample min-max。\n          训练阶段用“训练集全局 min/max”统一缩放）\n输出  : 光谱均值 μ̂ — 形状 **(batch, 1)**\n\nMC Dropout : 在推理阶段保持 model.train()，多次前向取均值/方差。\n\"\"\"\n\n# from __future__ import annotations\nfrom typing import Tuple\nimport torch\nfrom torch import nn\n\n# 模型本体\nclass Net1D(nn.Module):\n    \"\"\"\n    Simple 1-D CNN aligned with the starter Keras model.\n\n    Parameters\n    ----------\n    dropout_rate : float\n        Dropout 概率（默认 0.2）\n\n    Forward\n    -------\n    x : torch.Tensor, shape (B, 1, 187)\n        White-light curve（已在训练阶段做统一归一化）\n    Returns\n    -------\n    torch.Tensor, shape (B, 1)\n        Predicted mean transit depth (linear)\n    Notes\n    -----\n    - Conv1d(kernel=3, padding=0, 'valid') + MaxPool1d(2) × 4\n    - 长度演化：187→92→45→21→9\n    - Flatten 维度 = 256 * 9 = 2304（与 Keras 对齐）\n    - 在第一层池化后加 BatchNorm1d(32)\n    \"\"\"\n    def __init__(self, dropout_rate:float=0.2)->None:\n        super().__init__()\n\n        self.features = nn.Sequential(\n            nn.Conv1d(1, 32, kernel_size=3), nn.ReLU(), # 187->185\n            nn.MaxPool1d(2), # 185->92\n            nn.BatchNorm1d(32),\n\n            nn.Conv1d(32, 64, 3), nn.ReLU(), # 92->90\n            nn.MaxPool1d(2), # 90->45\n\n            nn.Conv1d(64, 128, 3), nn.ReLU(), # 45->43\n            nn.MaxPool1d(2), # 43->21\n\n            nn.Conv1d(128, 256, 3), nn.ReLU(), # 21->`19\n            nn.MaxPool1d(2), # 19->9\n        )\n        with torch.no_grad():\n            L = self.features(torch.zeros(1,1,187)).shape[-1]\n        self.head = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(256 * L, 500), nn.ReLU(),\n            nn.Dropout(dropout_rate),\n            nn.Linear(500, 100), nn.ReLU(),\n            nn.Dropout(dropout_rate/2),\n            nn.Linear(100, 1)  # 输出 μ̂\n        )\n\n    def forward(self, x:torch.Tensor)->torch.Tensor:\n        \"\"\"\n        Forward pass.\n            Parameters\n                ----------\n                x : torch.Tensor, shape (B, 1, 187)\n                    Normalised white-light curve.\n            Returns\n                -------\n                torch.Tensor, shape (B, 1)\n                    Mean transit depth.\n        \"\"\"\n        z = self.features(x)\n        return self.head(z)\n\n# Training / inference helpers\ndef train_1d_epoch(\n        model: nn.Module,\n        dataloader: torch.utils.data.DataLoader,\n        optim: torch.optim.Optimizer,\n        device:torch.device\n)->float:\n    \"\"\"\n        单个 epoch 训练函数。\n        Parameters\n        ----------\n        model      : Net1D (already on device)\n        dataloader : yields dict with keys \"wc\", \"y_mean\"\n        optim      : torch optimizer\n        device     : torch.device\n        Returns\n        -------\n        float : epoch mean loss (MSE)\n    \"\"\"\n    model.train()\n    mse = nn.MSELoss()\n    running = 0.0\n    for batch in dataloader:\n        # 注意：在 Notebook/主程序里，先对 batch[\"wc\"] 做“训练集全局 min/max”的统一归一化\n        x = batch[\"wc\"].to(device, non_blocking=True) #(B,1,187)\n        y = batch[\"y_mean\"].to(device, non_blocking=True)#(B,1)\n        optim.zero_grad()\n        y_hat = model(x)\n        loss = mse(y_hat, y)\n        loss.backward()\n        optim.step()\n        running += loss.item() * x.size(0)\n    return running / len(dataloader.dataset)\n\ndef enable_dropout_only(model:nn.Module)->None:\n    \"\"\"\n        Keep model in eval() but force ONLY Dropout layers into train() so MC dropout works,\n        while BatchNorm stays in eval() (frozen running stats).\n    \"\"\"\n    for mod in model.modules():\n        if isinstance(mod, (torch.nn.Dropout,torch.nn.AlphaDropout)):\n            mod.train()\n        elif isinstance(mod, (torch.nn.BatchNorm1d, torch.nn.BatchNorm2d)):\n            mod.eval()\n\n@torch.no_grad()\ndef predict_mc_dropout(\n        model:nn.Module,\n        x:torch.Tensor,\n        n_samples:int=200,\n)->Tuple[torch.Tensor, torch.Tensor]:\n    \"\"\"\n        多次前向推理以估计 (μ̂, σ_μ̂)。\n        Parameters\n        ----------\n        model      : Net1D (keep in *train* mode to enable Dropout)\n        x          : torch.Tensor, shape (B, 1, 374)  — already on same device\n        n_samples  : int, default 200， 采样次数\n        Returns\n        -------\n        μ_hat : torch.Tensor, shape (B, 1)\n        σ_hat : torch.Tensor, shape (B, 1)\n    \"\"\"\n    model.eval()\n    enable_dropout_only(model)\n\n    preds = []\n    for _ in range(n_samples):\n        preds.append(model(x))\n    preds = torch.stack(preds, dim=0)\n    mu = preds.mean(dim=0)\n    sig = preds.std(dim=0)\n    return mu, sig\n\n@torch.no_grad()\ndef mc_1d_on_loader(\n    model: nn.Module,\n    dataloader: torch.utils.data.DataLoader,\n    device: torch.device,\n    n_samples: int = 200,\n):\n    \"\"\"\n    返回:\n      mu  : (N,)   —— 反归一化后的 y_mean 预测均值\n      sig : (N,)   —— 反归一化后的 y_mean 预测标准差\n      y   : (N,)   —— 真值（未归一化）\n    依赖外部的 scale/unscale 函数: scale_wc_batch, unscale_y_batch 以及全局 ymin/ymax\n    \"\"\"\n    model.eval()\n    enable_dropout_only(model)\n\n    mu_list, sig_list, y_true = [], [], []\n    for batch in dataloader:\n        x = scale_wc_batch(batch[\"wc\"].to(device))          # (B,1,187) in [0,1]\n        y = batch[\"y_mean\"].to(device)                      # (B,1) original scale\n\n        # 多次前向，收集在 CPU，避免显存积累\n        outs = []\n        for _ in range(n_samples):\n            outs.append(model(x).detach().cpu())            # (B,1) in [0,1] scale\n        outs = torch.stack(outs, dim=0)                     # (S,B,1)\n\n        mu_norm = outs.mean(dim=0).squeeze(-1)              # (B,)\n        sig_norm = outs.std(dim=0).squeeze(-1)              # (B,)\n\n        # 反归一化：均值线性可直接整体反归一化；std 乘以尺度\n        mu = unscale_y_batch(mu_norm.numpy().reshape(-1, 1)).squeeze(1)                 # (B,)\n        sig = (sig_norm.numpy() * (ymax - ymin + 1e-8))                                   # (B,)\n\n        mu_list.append(mu)\n        sig_list.append(sig)\n        y_true.append(y.cpu().numpy().squeeze(1))\n\n    mu = np.concatenate(mu_list, axis=0)    # (N,)\n    sig = np.concatenate(sig_list, axis=0)  # (N,)\n    y  = np.concatenate(y_true, axis=0)     # (N,)\n    return mu, sig, y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T18:29:20.995161Z","iopub.execute_input":"2025-08-26T18:29:20.995396Z","iopub.status.idle":"2025-08-26T18:29:21.014905Z","shell.execute_reply.started":"2025-08-26T18:29:20.995371Z","shell.execute_reply":"2025-08-26T18:29:21.014322Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **2D CNN**# ","metadata":{}},{"cell_type":"code","source":"# net2d.py\n# --------------------------------------------\n# 2D CNN (PyTorch) for Ariel ADC 2025\n# 输入: map2d ∈ R^{B, 1, 40, 283}  —— AIRS 2D切片，已做星光谱归一化 & 整块去均值\n# 目标: y_shift ∈ R^{B, 283}       —— 去均值后的光谱（均值≈0）\n# 规范化: 训练阶段用 train 统计量做 [-1,1] 归一化：\n#        map2d / data_abs_max,  y_shift / targets_abs_max\n# 输出: Δ̂ ∈ R^{B, 283}（线性）\n# 不确定度: 用 MC-Dropout 多次前向（保持 model.train()）估计均值与方差\n# --------------------------------------------\nfrom typing import Tuple\nimport torch\nimport torch.nn as nn\nimport numpy as np\n\n\ndef _same_pad_3x1()->Tuple[int, int]:\n    \"\"\"\n    返回 3x1 卷积的 padding='same'\n    \"\"\"\n    return (1,0)\n\ndef _same_pad_1x3()->Tuple[int, int]:\n    \"\"\"\n    返回 1x3 卷积的 padding='same'\n    \"\"\"\n    return (0,1)\n\n# 2D-CNN（预测去均值的 283 维形状 Δ̂）\n# 简单复刻 Keras 的拓扑（时向和波向交替卷积/池化）。用动态推断 Flatten 维度，不纠结手算。\n\nclass Net2D(nn.Module):\n    \"\"\"\n    2D CNN（贴近原作者 Keras 拓扑，语义化命名 + 动态展平）\n    Parameters\n    ----------\n    time_len : int\n        输入的时间维长度（默认 40，对应固定窗口 75~115）\n    wave_len : int\n        输入的波长维长度（默认 283，对应 39..321）\n    dropout : float\n        全连接头部的 dropout 概率\n    Forward\n    -------\n    x : torch.Tensor, shape (B, 1, time_len, wave_len)\n        2D 观测（map2d），训练时请先做全局 [-1,1] 规范化：x /= data_abs_max\n    Returns\n    -------\n    torch.Tensor, shape (B, wave_len)\n        预测的 Δ̂（零均值的光谱形状）\n    \"\"\"\n    def __init__(self,time_len:int=40,wave_len:int=283,dropout:float=0.2):\n        super().__init__()\n        self.time_len = time_len\n        self.wave_len = wave_len\n\n        # 时向卷积块（与原作者的(3,1)卷积一致）\n        self.block_t1 = nn.Sequential(\n            nn.Conv2d(1, 32, kernel_size=(3, 1), padding=_same_pad_3x1()),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=(2, 1)), # 时间维/2\n            nn.BatchNorm2d(32)\n        )\n        self.block_t2 = nn.Sequential(\n            nn.Conv2d(32, 64, kernel_size=(3, 1), padding=_same_pad_3x1()),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=(2, 1)),\n        )\n        self.block_t3 = nn.Sequential(\n            nn.Conv2d(64, 128, kernel_size=(3, 1), padding=_same_pad_3x1()),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=(2, 1)),\n        )\n        self.block_t4 = nn.Sequential(\n            nn.Conv2d(128, 256, kernel_size=(3,1), padding=_same_pad_3x1()),\n            nn.ReLU(),\n        )\n\n        # 波向卷积块（与原作者的(1,3)卷积一致）\n        self.block_w1 = nn.Sequential(\n            nn.Conv2d(256, 32, kernel_size=(1,3), padding=_same_pad_1x3()),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=(1,2)), # 波长维/2\n            nn.BatchNorm2d(32)\n        )\n        self.block_w2 = nn.Sequential(\n            nn.Conv2d(32, 64, kernel_size=(1, 3), padding=_same_pad_1x3()),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=(1, 2)),\n        )\n        self.block_w3 = nn.Sequential(\n            nn.Conv2d(64, 128, kernel_size=(1, 3), padding=_same_pad_1x3()),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=(1, 2)),\n        )\n        self.block_w4 = nn.Sequential(\n            nn.Conv2d(128, 256, kernel_size=(1, 3), padding=_same_pad_1x3()),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=(1, 2)),\n        )\n\n        # 动态推断展平维度，避免Linear维度不匹配\n        with torch.no_grad():\n            dummy = torch.zeros(1,1,time_len,wave_len)\n            z = self._forward_features(dummy)\n            self.flat_dim = z.numel()\n\n        # 全连接头（与原作者700-》283对齐）\n        self.head = nn.Sequential(\n            nn.Flatten(),\n            nn.LazyLinear(700), nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(700, wave_len),\n        )\n\n    def _forward_features(self, x: torch.Tensor)->torch.Tensor:\n        \"\"\"仅做卷积与池化，返回 (B, C, T', W')。\"\"\"\n        z = self.block_t1(x)  # (B, 32, 20, 283)\n        z = self.block_t2(z)  # (B, 64, 10, 283)\n        z = self.block_t3(z)  # (B,128,  5, 283)\n        z = self.block_t4(z)  # (B,256,  5, 283)\n        z = self.block_w1(z)  # (B, 32,  5, ~141)\n        z = self.block_w2(z)  # (B, 64,  5, ~70)\n        z = self.block_w3(z)  # (B,128,  5, ~35)\n        z = self.block_w4(z)  # (B,256,  5, ~17)\n        return z\n\n    def forward(self, x:torch.Tensor)->torch.Tensor:\n        \"\"\"\n            Parameters\n            ----------\n                x : torch.Tensor, shape (B, 1, time_len, wave_len)\n                    训练阶段请先做：x = x / data_abs_max\n            Returns\n            -------\n                torch.Tensor, shape (B, wave_len)\n                    预测 Δ̂\n        \"\"\"\n        z = self._forward_features(x)\n        y_hat = self.head(z) # (B, 283)\n        return y_hat\n\n# ---------------------------\n# 训练 / 评估 / MC Dropout API\n# ---------------------------\n\ndef train_2d_epoch(\n        model:nn.Module, dataloader:torch.utils.data.DataLoader,\n        optimizer: torch.optim.Optimizer,\n        device:torch.device,\n        data_abs_max:float,\n        targets_abs_max:float,\n)->float:\n    \"\"\"\n        训练单个 epoch（MSE 对 y_shift）\n        Parameters\n        ----------\n        model : Net2D（已 to(device)）\n        dataloader : 产出 dict，包括 \"map2d\", \"y_shift\"\n            - batch[\"map2d\"] : (B,1,40,283)\n            - batch[\"y_shift\"]: (B,283)\n        optimizer : torch optimizer\n        device : torch.device\n        data_abs_max : float\n            训练集上 map2d 的绝对最大值（全局），用于 map2d 归一化\n        targets_abs_max : float\n            训练集上 y_shift 的绝对最大值（全局），用于 y_shift 归一化\n        Returns\n        -------\n        float : epoch 平均损失（MSE）\n    \"\"\"\n    model.train()\n    mse = nn.MSELoss()\n    running = 0.0\n\n    for batch in dataloader:\n        # [-1,1] 归一化（按训练集统计量）\n        x = (batch[\"map2d\"]/(data_abs_max+1e-12)).to(device, non_blocking=True) # (B,1,40,283)\n        y = (batch[\"y_shift\"]/(targets_abs_max+1e-12)).to(device, non_blocking=True) # (B,283)\n\n        optimizer.zero_grad(set_to_none=True)\n        y_hat = model(x)\n        loss = mse(y_hat,y)\n        loss.backward()\n        optimizer.step()\n        running += loss.item() * x.size(0)\n\n    return running/len(dataloader.dataset)\n\n@torch.no_grad()\ndef eval_2d_rmse(\n        model:nn.Module,\n        dataloader:torch.utils.data.DataLoader,\n        device:torch.device,\n        data_abs_max:float,\n        targets_abs_max:float,\n) -> Tuple[float, np.ndarray, np.ndarray]:\n    \"\"\"\n        评估 RMSE（反归一化后，以物理尺度）\n        Returns\n        -------\n        rmse : float\n            预测 Δ̂ 与真值 Δ 的 RMSE\n        preds : np.ndarray, shape (N, 283)\n            反归一化后的预测\n        trues : np.ndarray, shape (N, 283)\n            真值（未归一化）\n    \"\"\"\n    model.eval()\n    preds, trues = [], []\n    for batch in dataloader:\n        x = (batch[\"map2d\"]/(data_abs_max+1e-12)).to(device)\n        y = batch[\"y_shift\"].to(device)\n        y_hat = model(x).cpu().numpy() * (targets_abs_max+1e-12) # 反归一化\n        preds.append(y_hat)\n        trues.append(y.cpu().numpy())\n\n    preds = np.concatenate(preds, axis=0) # (N,283)\n    trues = np.concatenate(trues, axis=0) # (N,283)\n    rmse = float(((preds-trues)**2).mean()**0.5)\n\n    return rmse, preds, trues\n\ndef enable_dropout_only(model:nn.Module)->None:\n    \"\"\"\n        Keep model in eval() but force ONLY Dropout layers into train() so MC dropout works,\n        while BatchNorm stays in eval() (frozen running stats).\n    \"\"\"\n    for mod in model.modules():\n        if isinstance(mod, (torch.nn.Dropout,torch.nn.AlphaDropout)):\n            mod.train()\n        elif isinstance(mod, (torch.nn.BatchNorm1d, torch.nn.BatchNorm2d)):\n            mod.eval()\n\n@torch.no_grad()\ndef mc_2d_on_loader(\n        model:nn.Module,\n        dataloader:torch.utils.data.DataLoader,\n        device:torch.device,\n        n_samples:int,\n        data_abs_max:float,\n        targets_abs_max:float,\n):\n    \"\"\"\n        MC Dropout：重复前向 n 次（保持 model.train() 使 Dropout 生效）估计 (μ, σ)\n        Returns\n        -------\n        mu : np.ndarray, shape (N,283)\n            预测均值（反归一化）\n        sig: np.ndarray, shape (N,283)\n            预测标准差（反归一化）\n        y_true : np.ndarray, shape (N,283)\n            真值（未归一化）\n    \"\"\"\n    model.eval()\n    enable_dropout_only(model)\n    mu_list, sig_list, y_true = [],[],[]\n\n    for batch in dataloader:\n        x = (batch[\"map2d\"]/(data_abs_max+1e-12)).to(device)\n        y = batch[\"y_shift\"].to(device)\n\n        samples = []\n        for _ in range(n_samples):\n            samples.append(model(x).cpu().numpy()) #(B,283)\n        samples =  np.stack(samples,axis=0) # (S,B,283)\n\n        mu = samples.mean(axis=0)*(targets_abs_max+1e-12)\n        sig = samples.std(axis=0)*(targets_abs_max+1e-12)\n\n        mu_list.append(mu)\n        sig_list.append(sig)\n        y_true.append(y.cpu().numpy())\n\n    mu = np.concatenate(mu_list,axis=0)\n    sig = np.concatenate(sig_list,axis=0)\n    y_true = np.concatenate(y_true,axis=0)\n    return mu, sig, y_true","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T18:29:21.016674Z","iopub.execute_input":"2025-08-26T18:29:21.016844Z","iopub.status.idle":"2025-08-26T18:29:21.03887Z","shell.execute_reply.started":"2025-08-26T18:29:21.016831Z","shell.execute_reply":"2025-08-26T18:29:21.038373Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, time, random, math\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n\n# from dataset_cnn import ArielCNNDataset, resample_2d_time, AIR_SLICE, CUT_BEGIN, CUT_END\n# from net1d import Net1D, train_1d_epoch, predict_mc_dropout\n\nBASE = \"/kaggle/input/ariel-data-challenge-2025\"\nOLD_CACHE = \"cnn_cache/train\"\nNEW_CACHE = \"cnn_cache_wc187/train\"\nos.makedirs(NEW_CACHE,exist_ok=True)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T18:29:21.039456Z","iopub.execute_input":"2025-08-26T18:29:21.039617Z","iopub.status.idle":"2025-08-26T18:29:21.139434Z","shell.execute_reply.started":"2025-08-26T18:29:21.039603Z","shell.execute_reply":"2025-08-26T18:29:21.138661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rebuild_wc_cache_to_187(base_dir, old_cache, new_cache):\n    \"\"\"\n    只重算 wc -> (1,187),之前的是错误拼接导致的(1,374)，map2d/y_* 尽量从 old_cache 拷贝；old_cache 没有就现算。\n    \"\"\"\n    ds = ArielCNNDataset(base_dir, split=\"train\", return_labels=True, cache_calib=True, disk_cache_dir=None)\n    \n    def build_wc187_from_airs(airs_2d: torch.Tensor)->np.ndarray:\n        # airs_2d: (5625,283)  -> time bin -> (187,283) -> white_curve\n        airs_2d_bin = resample_2d_time(airs_2d)                  # (187,283)\n        wc_mean = airs_2d_bin.mean()\n        wc_187  = airs_2d_bin.sum(dim=1) / (wc_mean + 1e-8)     # (187,)\n        return wc_187.unsqueeze(0).numpy().astype(\"float32\")     # (1,187)\n    \n    print(\"Rebuilding wc cache to 187 points... ->\", os.path.abspath(new_cache))\n    for pid in tqdm(ds.pids):\n        outp = os.path.join(new_cache, f\"{pid}.npz\")\n        if os.path.exists(outp):\n            continue  # 已存在，跳过\n        \n        # 读取旧缓存(拷贝map2d/y_*)\n        oldp = os.path.join(old_cache, f\"{pid}.npz\")\n        z_old = np.load(oldp) if os.path.exists(oldp) else None\n        \n        # 读AIRS->wc\n        airs = ds._load_sensor(pid, \"AIRS-CH0\")  # (5625, 32,356)\n        airs_2d = airs.sum(dim=1)[:, AIR_SLICE[0]:AIR_SLICE[1]+1].cpu()  # (5625,283)\n        wc = build_wc187_from_airs(airs_2d)  # (1,187)\n        \n        if z_old is not None:\n            np.savez_compressed(\n                outp,\n                wc=wc,\n                map2d=z_old[\"map2d\"],  # (1,40,283)\n                y_mean=z_old[\"y_mean\"],  # (283,)\n                y_shift=z_old[\"y_shift\"]  # (283,)\n            )\n        else:\n            # 没有旧缓存，现算map2d & y\n            airs_2d_bin = resample_2d_time(airs_2d)                      # (187,283)\n\n            # 星光谱归一化 -> in-transit 切片 -> 整块去均值（对齐作者）\n            oot_left  = airs_2d_bin[:50].mean(dim=0)\n            oot_right = airs_2d_bin[-50:].mean(dim=0)\n            star_spec = oot_left + oot_right\n            airs_2d_norm = airs_2d_bin / (star_spec.clamp_min(1e-8))\n            map_slice = airs_2d_norm[CUT_BEGIN:CUT_END]                  # (40,283)\n            map2d = (map_slice - map_slice.mean()).unsqueeze(0).numpy().astype(\"float32\")  # (1,40,283)\n\n            # y 标签\n            y_full = torch.from_numpy(ds._y_map[str(pid)]).float()\n            y_mean = y_full.mean().unsqueeze(0).numpy().astype(\"float32\")\n            y_shift = (y_full - y_mean).numpy().astype(\"float32\")\n            np.savez_compressed(outp, wc=wc, map2d=map2d, y_mean=y_mean, y_shift=y_shift)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T18:29:21.140222Z","iopub.execute_input":"2025-08-26T18:29:21.140448Z","iopub.status.idle":"2025-08-26T18:29:21.15262Z","shell.execute_reply.started":"2025-08-26T18:29:21.140421Z","shell.execute_reply":"2025-08-26T18:29:21.152025Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rebuild_wc_cache_to_187(BASE, OLD_CACHE, NEW_CACHE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T18:29:21.153376Z","iopub.execute_input":"2025-08-26T18:29:21.15421Z","iopub.status.idle":"2025-08-26T19:15:39.608336Z","shell.execute_reply.started":"2025-08-26T18:29:21.154186Z","shell.execute_reply":"2025-08-26T19:15:39.607698Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 8:2 划分训练和验证集\ndef split_planet_ids(planet_ids, train_ratio=0.8, seed=42):\n    rng = random.Random(seed)\n    ids = list(map(str,planet_ids))\n    rng.shuffle(ids)\n    k = int(len(ids) * train_ratio)\n    return ids[:k], ids[k:]  # 返回训练集和验证集的 planet_id 列表\n\n#用 NEW_CACHE里已有的pid\npid_files = sorted(glob.glob(os.path.join(NEW_CACHE, \"*.npz\")))\nall_ids = [os.path.splitext(os.path.basename(p))[0]for p in pid_files]\ntrain_ids, valid_ids = split_planet_ids(all_ids, train_ratio=0.8, seed=42)\nlen(all_ids), len(train_ids), len(valid_ids)  # 检查划分结果","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:15:39.609358Z","iopub.execute_input":"2025-08-26T19:15:39.609599Z","iopub.status.idle":"2025-08-26T19:15:39.621715Z","shell.execute_reply.started":"2025-08-26T19:15:39.609579Z","shell.execute_reply":"2025-08-26T19:15:39.62118Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\n统计全局归一化参数（对齐作者方式）\n1D：wc 用训练集 全局 min/max；y_mean 用训练集 min/max。\n2D：targets（y_shift）和 obs（map2d）都用训练集的绝对最大值进行 [-1,1] 规范化。\n'''\n# 扫描train缓存，统计scaler\nW_list, Ymean_list, Yshift_list, MAP_list = [], [], [], []\n\nfor pid in tqdm(train_ids, desc=\"scan train scalers\"):\n    z = np.load(os.path.join(NEW_CACHE, f\"{pid}.npz\"))\n    W_list.append(z[\"wc\"].squeeze(0))  # (187,)\n    Ymean_list.append(z[\"y_mean\"].squeeze(0))  # ()\n    Yshift_list.append(z[\"y_shift\"])  # (283,)\n    MAP_list.append(z[\"map2d\"].squeeze(0))  # (40,283)\nW = np.stack(W_list) # (N_train, 187)\nYmean = np.stack(Ymean_list)  # (N_train, 283)\nYshift = np.stack(Yshift_list)  # (N_train, 283)\nMAP = np.stack(MAP_list)  # (N_train, 40, 283)\n\n# 计算全局 min/max\n# 1D scalers\nwc_min, wc_max = W.min(), W.max()\ny_min, y_max = Ymean.min(), Ymean.max()\n\n# 2D scalers\ntargets_abs_max = max(abs(Yshift.min()), abs(Yshift.max()))\ndata_abs_max = max(abs(MAP.min()), abs(MAP.max()))\n\nprint(\"wc_min/max:\", float(wc_min), float(wc_max))\nprint(\"y_mean min/max:\", float(y_min), float(y_max))\nprint(\"targets_abs_max:\", float(targets_abs_max))\nprint(\"data_abs_max:\", float(data_abs_max))\n\n# 保存scalers,训练/预测都用同一份\nos.makedirs(\"models\", exist_ok=True)\nnp.savez(\"models/scalers_v1.npz\",\n         wc_min=wc_min, wc_max=wc_max,\n         y_min=y_min, y_max=y_max,\n         targets_abs_max=targets_abs_max,\n         data_abs_max=data_abs_max)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:15:39.622387Z","iopub.execute_input":"2025-08-26T19:15:39.622613Z","iopub.status.idle":"2025-08-26T19:15:40.813263Z","shell.execute_reply.started":"2025-08-26T19:15:39.622592Z","shell.execute_reply":"2025-08-26T19:15:40.812513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bs = 32\ncuda = device.type == \"cuda\"\ndl_kwargs = dict(batch_size=bs, num_workers=0, pin_memory=cuda, persistent_workers=False)\nds_train = ArielCNNDataset(BASE, planet_ids=train_ids, split=\"train\",\n                           return_labels=True, disk_cache_dir=NEW_CACHE)\nds_valid = ArielCNNDataset(BASE, planet_ids=valid_ids, split=\"train\",\n                           return_labels=True, disk_cache_dir=NEW_CACHE)\ndl_train = DataLoader(ds_train, shuffle=True, **dl_kwargs)\ndl_valid = DataLoader(ds_valid, shuffle=False, **dl_kwargs)\n\nb = next(iter(dl_train))\nprint(\"Batch shapes:\", b[\"wc\"].shape, b[\"map2d\"].shape, b[\"y_mean\"].shape, b[\"y_shift\"].shape) # (B,1,187), (B,1,40,283), (B,1), (B,283)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:15:40.815516Z","iopub.execute_input":"2025-08-26T19:15:40.815724Z","iopub.status.idle":"2025-08-26T19:15:41.133416Z","shell.execute_reply.started":"2025-08-26T19:15:40.815707Z","shell.execute_reply":"2025-08-26T19:15:41.132643Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 训练辅助：归一化/反归一化函数\n# 载入scalers\nsc = np.load(\"models/scalers_v1.npz\")\nwc_min = float(sc[\"wc_min\"])\nwc_max = float(sc[\"wc_max\"])\nymin = float(sc[\"y_min\"])\nymax = float(sc[\"y_max\"])\ntargets_abs_max = float(sc[\"targets_abs_max\"])\ndata_abs_max = float(sc[\"data_abs_max\"])\n\n# 1D helpers\ndef scale_wc_batch(x): # x: (B,1,187)\n    return (x-wc_min) / (wc_max - wc_min + 1e-8)  # scale to [0,1]\ndef scale_y_batch(y): # y: (B,1)\n    return (y - ymin) / (ymax - ymin + 1e-8)  # scale to [0,1]\ndef unscale_y_batch(y): # x: (B,1)\n    return y * (ymax - ymin + 1e-8) + ymin  # unscale to original\n\n# 2D helpers\ndef scale_obs2d_batch(m): # m: (B,1,40,283)\n    return m / (data_abs_max + 1e-8)  # scale to [-1,1]\ndef scale_target2d_batch(t): # t: (B,283)\n    return t / (targets_abs_max + 1e-8)  # scale to [-1,1]\ndef unscale_target2d_batch(t): # t: (B,283)\n    return t * (targets_abs_max + 1e-8)  # unscale to original","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:15:41.134224Z","iopub.execute_input":"2025-08-26T19:15:41.134452Z","iopub.status.idle":"2025-08-26T19:15:41.141971Z","shell.execute_reply.started":"2025-08-26T19:15:41.134428Z","shell.execute_reply":"2025-08-26T19:15:41.141022Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1D CNN 模型\nnet1d = Net1D(dropout_rate=0.2).to(device)\nopt = torch.optim.Adam(net1d.parameters(), lr=0.001, weight_decay=1e-5)\nsch = torch.optim.lr_scheduler.StepLR(opt, step_size=200, gamma=0.2)\n\ndef train_1d_epoch_scaled(model, dataloader, optim, device):\n    model.train()\n    mse = torch.nn.MSELoss()\n    running = 0.0\n    for batch in dataloader:\n        x = scale_wc_batch(batch[\"wc\"].to(device, non_blocking=True))\n        y = scale_y_batch(batch[\"y_mean\"].to(device, non_blocking=True))\n        optim.zero_grad()\n        y_hat = model(x)\n        loss = mse(y_hat, y)\n        loss.backward()\n        optim.step()\n        running += loss.item() * x.size(0)\n    return running / len(dataloader.dataset)\n\n@torch.no_grad()\ndef eval_1d_rmse(model, dataloader, device):\n    model.eval()\n    preds, trues = [], []\n    for batch in dataloader:\n        x = scale_wc_batch(batch[\"wc\"].to(device))\n        y = batch[\"y_mean\"].to(device)\n        p = unscale_y_batch(model(x).cpu().numpy())\n        preds.append(p[:,0])\n        trues.append(y.cpu().numpy()[:,0])\n    preds = np.concatenate(preds); trues = np.concatenate(trues)\n    rmse = np.sqrt(np.mean((preds - trues) ** 2))\n    return rmse, preds, trues","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:15:41.142808Z","iopub.execute_input":"2025-08-26T19:15:41.143088Z","iopub.status.idle":"2025-08-26T19:15:43.967657Z","shell.execute_reply.started":"2025-08-26T19:15:41.143065Z","shell.execute_reply":"2025-08-26T19:15:43.966832Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS_1D = 300\nbest_rmse = 1e9; best_path = \"models/net1d.pt\"\nt0 = time.time()\nfor ep in range(1, EPOCHS_1D+1):\n    loss = train_1d_epoch_scaled(net1d, dl_train, opt,device)\n    sch.step()\n    if ep%10 == 0 or ep == 1:\n        rmse_tr, _, _= eval_1d_rmse(net1d, dl_train, device)\n        rmse_va, _, _ = eval_1d_rmse(net1d, dl_valid, device)\n        print(f\"[1D] ep {ep:4d} loss {loss:.6} RMSE(tr/va) {rmse_tr:.6e}/{rmse_va:.6e}\")\n        if rmse_va < best_rmse:\n            best_rmse = rmse_va\n            torch.save(net1d.state_dict(), best_path)\n            print(f\"New best RMSE: {best_rmse:.6e}, saved to {best_path}\")\nprint(\"1D done in %.1f s\"%(time.time()-t0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:15:43.968552Z","iopub.execute_input":"2025-08-26T19:15:43.96888Z","iopub.status.idle":"2025-08-26T19:20:55.315544Z","shell.execute_reply.started":"2025-08-26T19:15:43.968862Z","shell.execute_reply":"2025-08-26T19:20:55.314834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 载入最佳 1D 并画散点\nnet1d.load_state_dict(torch.load(\"models/net1d.pt\", map_location=device))\nrmse_va, p_va, y_va = eval_1d_rmse(net1d, dl_valid, device)\nplt.figure(); plt.scatter(y_va, p_va, s=8, alpha=0.5)\nm1, m2 = min(y_va.min(), p_va.min()), max(y_va.max(), p_va.max())\nplt.plot([m1,m2],[m1,m2],'k--'); plt.title(f\"1D valid RMSE={rmse_va:.6e}\")\nplt.xlabel(\"y_mean true\"); plt.ylabel(\"y_mean pred\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:20:55.316445Z","iopub.execute_input":"2025-08-26T19:20:55.316674Z","iopub.status.idle":"2025-08-26T19:20:55.807927Z","shell.execute_reply.started":"2025-08-26T19:20:55.316656Z","shell.execute_reply":"2025-08-26T19:20:55.807267Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef mc_1d_on_loader(model, dataloader, device, n_samples=500):\n    model.eval()\n    mu_list, sig_list, y_true = [], [], []\n    for batch in dataloader:\n        x = scale_wc_batch(batch[\"wc\"].to(device))\n        y = batch[\"y_mean\"].to(device)\n        mu, sig = predict_mc_dropout(model, x, n_samples=n_samples)\n        mu = unscale_y_batch(mu).cpu().numpy()[:,0]\n        sig = sig.cpu().numpy()[:,0] * (ymax-ymin+1e-8)\n        mu_list.append(mu); sig_list.append(sig); y_true.append(y.cpu().numpy()[:,0])\n    return np.concatenate(mu_list), np.concatenate(sig_list), np.concatenate(y_true)\n\nmu_va, sig_va, y_va = mc_1d_on_loader(net1d, dl_valid, device, n_samples=200)\nprint(\"1D MC-Dropout: mu/std mean\", mu_va.mean(), sig_va.mean())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:20:55.808699Z","iopub.execute_input":"2025-08-26T19:20:55.808924Z","iopub.status.idle":"2025-08-26T19:20:57.071968Z","shell.execute_reply.started":"2025-08-26T19:20:55.808897Z","shell.execute_reply":"2025-08-26T19:20:57.071395Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mu_va, sig_va, y_va = mc_1d_on_loader(net1d, dl_valid, device, n_samples=200)\nnp.savez(\"models/one_d_valid.npz\", mu=mu_va, sig=sig_va, y=y_va)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:20:57.072621Z","iopub.execute_input":"2025-08-26T19:20:57.072812Z","iopub.status.idle":"2025-08-26T19:20:58.326614Z","shell.execute_reply.started":"2025-08-26T19:20:57.072794Z","shell.execute_reply":"2025-08-26T19:20:58.326053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from net2d import Net2D, train_2d_epoch, eval_2d_rmse, mc_2d_on_loader\n\nnet2d = Net2D(time_len=40, wave_len=283, dropout=0.2).to(device)\nopt2 = torch.optim.Adam(net2d.parameters(),lr=1e-3,weight_decay=1e-6)\n\nEPOCHS_2D = 120\nbest_rmse, best_path = 1e9, \"models/net2d.pt\"\nfor ep in range(1, EPOCHS_2D+1):\n    loss = train_2d_epoch(net2d,dl_train,opt2,device,data_abs_max,targets_abs_max)\n    if ep%10 == 0 or ep ==1:\n        rmse_tr,_,_ = eval_2d_rmse(net2d, dl_train, device, data_abs_max, targets_abs_max)\n        rmse_va, _, _ = eval_2d_rmse(net2d, dl_valid, device, data_abs_max, targets_abs_max)\n        print(f\"[2D] ep {ep:4d} loss {loss:.6} RMSE(tr/va) {rmse_tr:.6e}/{rmse_va:.6e}\")\n        if rmse_va < best_rmse:\n            best_rmse = rmse_va\n            os.makedirs(\"models\", exist_ok=True)\n            torch.save(net2d.state_dict(),best_path)\n            print(\"  saved:\", best_path)\n\n# MC Dropout 20 次\nnet2d.load_state_dict(torch.load(\"models/net2d.pt\", map_location=device))\nmu2_va, sig2_va, y2_va =mc_2d_on_loader(net2d,dl_valid,device,n_samples=20,data_abs_max=data_abs_max,targets_abs_max=targets_abs_max)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:20:58.327345Z","iopub.execute_input":"2025-08-26T19:20:58.327583Z","iopub.status.idle":"2025-08-26T19:25:15.162288Z","shell.execute_reply.started":"2025-08-26T19:20:58.327562Z","shell.execute_reply":"2025-08-26T19:25:15.161694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 加载一下训练完的结果\n# from net2d import Net2D, eval_2d_rmse\n\n# -- load best 2D --\nnet2d = Net2D(time_len=40, wave_len=283, dropout=0.2).to(device)\nnet2d.load_state_dict(torch.load(\"models/net2d.pt\", map_location=device))\n\nrmse_va, preds_va, trues_va = eval_2d_rmse(\n    net2d, dl_valid, device, data_abs_max, targets_abs_max\n)\nprint(f\"[2D] valid RMSE (Δ̂): {rmse_va:.6e}  ({rmse_va*1e6:.1f} ppm)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:25:15.162983Z","iopub.execute_input":"2025-08-26T19:25:15.16317Z","iopub.status.idle":"2025-08-26T19:25:15.554831Z","shell.execute_reply.started":"2025-08-26T19:25:15.163155Z","shell.execute_reply":"2025-08-26T19:25:15.554183Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# wavelengths\nwavelength = pd.read_csv(os.path.join(BASE, \"wavelengths.csv\")).values.squeeze()\n\nper_wl_rmse = np.sqrt(((preds_va - trues_va)**2).mean(axis=0))  # (283,)\nplt.figure(figsize=(8,3))\nplt.plot(wavelength, per_wl_rmse*1e6)\nplt.xlabel(\"Wavelength (μm)\")\nplt.ylabel(\"RMSE (ppm)\")\nplt.title(\"Per-wavelength RMSE on valid (2D Δ̂)\")\nplt.grid(alpha=0.3)\nplt.show()\n\n# 做 2D MC Dropout（例如 50 次）\nmu2_va, sig2_va, y2_va = mc_2d_on_loader(\n    net2d, dl_valid, device, n_samples=50,\n    data_abs_max=data_abs_max, targets_abs_max=targets_abs_max\n)\n\n# 画几个样本\nidxs = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10]  # 自己改\nfor i in idxs:\n    plt.figure(figsize=(7,3))\n    plt.title(f\"2D shape prediction (valid idx={i})\")\n    plt.plot(wavelength, y2_va[i], color=\"tomato\", label=\"Δ true\")\n    plt.plot(wavelength, mu2_va[i], \".k\", ms=3, label=\"Δ pred\")\n    plt.fill_between(wavelength, mu2_va[i]-sig2_va[i], mu2_va[i]+sig2_va[i],\n                     color=\"silver\", alpha=0.5, label=\"Δ ± σ\")\n    plt.xlabel(\"Wavelength (μm)\")\n    plt.ylabel(\"(Rp/Rs)^2 (zero-mean)\")\n    plt.legend()\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:25:15.555593Z","iopub.execute_input":"2025-08-26T19:25:15.555842Z","iopub.status.idle":"2025-08-26T19:25:22.85162Z","shell.execute_reply.started":"2025-08-26T19:25:15.555819Z","shell.execute_reply":"2025-08-26T19:25:22.850852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compose_final_spectrum(\n        mu_1d:np.ndarray, # (N,),   1D y_mean predictions (original scale)\n        sigma_1d:np.ndarray, # (N,),   1D std (original scale)\n        mu_2d:np.ndarray, # (N,283), 2D delta predictions (original scale)\n        sigma_2d: np.ndarray # (N,283), 2D std (original scale)\n)->tuple[np.ndarray,np.ndarray]:\n    \"\"\"\n    Combine μ̂ (scalar per sample) and Δ̂(λ) to get FINAL spectrum and its uncertainty.\n    Returns\n    -------\n    y_final_pred : (N, 283)\n        Final spectrum prediction for each sample.\n    sigma_final  : (N, 283)\n        Uncalibrated predictive std under independence assumption:\n        Var_final = Var_mu + Var_delta.\n    Notes\n    -----\n    - We assume independence between μ̂ and Δ̂ when combining uncertainties.\n    \"\"\"\n    y_final_pred = mu_1d[:,None] + mu_2d\n    sigma_final = np.sqrt(sigma_1d[:,None]**2 + sigma_2d**2)\n    return y_final_pred, sigma_final\n\ndef calibrate_uncertainty_scalar(\n        y_true: np.ndarray, # (N,283)\n        y_pred: np.ndarray, # (N,283)\n        sigma_pred: np.ndarray # (N,283)\n)->tuple[float,float]:\n    \"\"\"\n    Calibrate predictive std with a single scalar factor c to improve coverage.\n    Returns\n    -------\n    c : float\n        Scalar calibration factor. sigma <- c * sigma\n    rmse : float\n        RMSE on validation under original (physical) scale.\n    Notes\n    -----\n    - c = RMSE / mean(sigma_pred). This is a simple, robust scalar calibration.\n    \"\"\"\n    residual = y_pred - y_true\n    rmse = float(np.sqrt(np.mean(residual**2)))\n    denom = float(np.mean(sigma_pred)+1e-12)\n    c = rmse/denom\n    return c, rmse\n\ndef compute_coverage(\n        y_true: np.ndarray, #(N,283)\n        y_pred:np.ndarray, #(N,283)\n        sigma_pred:np.ndarray, #(N,283)\n        k: float = 1.0 # 1σ = 68%, 2σ = 95%\n)->float:\n    \"\"\"\n    Compute empirical coverage: fraction of points inside [y_pred ± k*sigma].\n    Returns\n    -------\n    coverage : float\n        Mean fraction of wavelengths covered over all samples.\n    \"\"\"\n    lower = y_pred - k*sigma_pred\n    upper = y_pred + k*sigma_pred\n    inside = (y_true>=lower) & (y_true<=upper)\n    return float(np.mean(inside))\n\ndef plot_final_sample(\n    wavelengths: np.ndarray,      # (283,)\n    y_true: np.ndarray,           # (283,)\n    y_pred: np.ndarray,           # (283,)\n    sigma: np.ndarray,            # (283,)\n    sample_id: str | int = \"\"\n) -> None:\n    \"\"\"\n    Plot final spectrum for a single sample with ±1σ band.\n    Notes\n    -----\n    - One figure per sample to keep visuals clean.\n    \"\"\"\n    plt.figure(figsize=(8, 3))\n    title = f\"FINAL spectrum (sample {sample_id})\" if sample_id != \"\" else \"FINAL spectrum\"\n    plt.title(title)\n    plt.plot(wavelengths, y_true, label=\"Target\")\n    plt.plot(wavelengths, y_pred, \".\", ms=3, label=\"Prediction\")\n    plt.fill_between(wavelengths, y_pred - sigma, y_pred + sigma, alpha=0.35, label=\"±1σ\")\n    plt.xlabel(\"Wavelength (μm)\")\n    plt.ylabel(r\"$(R_p/R_s)^2$\")\n    plt.legend()\n    plt.tight_layout()\n    plt.show()\n    \ndef save_valid_ouputs(\n        path: str,\n        y_true:np.ndarray, #(N,283)\n        y_pred: np.ndarray, #(N,283)\n        sigma_calibrated: np.ndarray, #(N,283)\n        mu_1d: np.ndarray, #(N,)\n        sigma_1d: np.ndarray, #(N,)\n        mu_2d: np.ndarray, #(N,283)\n        sigma_2d: np.ndarray\n)->None:\n    \"\"\"\n    Save arrays for later analysis / plotting.\n    Outputs\n    -------\n    NPZ file with keys:\n    - y_true, y_pred, sigma_calibrated, mu_1d, sigma_1d, mu_2d, sigma_2d\n    \"\"\"\n    np.savez_compressed(\n        path,\n        y_true=y_true,\n        y_pred=y_pred,\n        sigma_calibrated=sigma_calibrated,\n        mu_1d=mu_1d,\n        sigma_1d=sigma_1d,\n        mu_2d=mu_2d,\n        sigma_2d=sigma_2d,\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:25:22.852428Z","iopub.execute_input":"2025-08-26T19:25:22.852675Z","iopub.status.idle":"2025-08-26T19:25:22.864429Z","shell.execute_reply.started":"2025-08-26T19:25:22.852653Z","shell.execute_reply":"2025-08-26T19:25:22.863682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --------------------------\n# 1) Compose FINAL on valid\n#    (assumes you already have:\n#     mu_va, sig_va, y_va from 1D;\n#     mu2_va, sig2_va, y2_va from 2D;\n#     wavelength from BASE/wavelengths.csv)\n# --------------------------\n\n# Sanity checks on shapes\nassert mu_va.ndim == 1 and sig_va.ndim == 1, \"1D arrays must be (N,)\"\nassert mu2_va.ndim == 2 and sig2_va.ndim == 2, \"2D arrays must be (N,283)\"\nassert y_va.ndim == 1 and y2_va.ndim == 2, \"Truth arrays must be (N,) and (N,283)\"\nassert mu2_va.shape[1] == sig2_va.shape[1] == y2_va.shape[1] == 283, \"Wavelength dimension must be 283\"\nassert len(wavelength) == 283, \"wavelength vector must have length 283\"\n\n# Compose prediction and uncalibrated sigma\ny_final_pred, sigma_final = compose_final_spectrum(mu_va, sig_va, mu2_va, sig2_va)\n# Compose ground-truth final spectra (mean + shift)\ny_final_true = y_va[:, None] + y2_va\n\n# Compute calibration factor and RMSE (before calibration)\ncalibration_factor_c, rmse_final_before = calibrate_uncertainty_scalar(y_final_true, y_final_pred, sigma_final)\n\n# Calibrate sigma\nsigma_final_calibrated = sigma_final * calibration_factor_c\n\n# Coverage (after calibration; also report before)\ncoverage_68_before = compute_coverage(y_final_true, y_final_pred, sigma_final, k=1.0)\ncoverage_95_before = compute_coverage(y_final_true, y_final_pred, sigma_final, k=2.0)\ncoverage_68_after  = compute_coverage(y_final_true, y_final_pred, sigma_final_calibrated, k=1.0)\ncoverage_95_after  = compute_coverage(y_final_true, y_final_pred, sigma_final_calibrated, k=2.0)\n\n# Per-wavelength RMSE on FINAL prediction\nper_wl_rmse_final = np.sqrt(((y_final_pred - y_final_true) ** 2).mean(axis=0))  # (283,)\n\nprint(f\"[FINAL] RMSE (before calib): {rmse_final_before:.6e}\")\nprint(f\"[FINAL] Scalar calibration c: {calibration_factor_c:.4f}\")\nprint(f\"[FINAL] Coverage 68%  before/after: {coverage_68_before:.3f} / {coverage_68_after:.3f}\")\nprint(f\"[FINAL] Coverage 95%  before/after: {coverage_95_before:.3f} / {coverage_95_after:.3f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:25:22.865165Z","iopub.execute_input":"2025-08-26T19:25:22.865418Z","iopub.status.idle":"2025-08-26T19:25:22.883905Z","shell.execute_reply.started":"2025-08-26T19:25:22.865398Z","shell.execute_reply":"2025-08-26T19:25:22.883347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --------------------------\n# 2) Plots\n# --------------------------\n\n# Per-wavelength RMSE curve\nplt.figure(figsize=(8, 3))\nplt.title(\"Per-wavelength RMSE on VALID (FINAL)\")\nplt.plot(wavelength, per_wl_rmse_final * 1e6)\nplt.xlabel(\"Wavelength (μm)\")\nplt.ylabel(\"RMSE (ppm)\")\nplt.grid(alpha=0.3)\nplt.tight_layout()\nplt.show()\n\n# Show a few samples\nsample_indices = [0, 1, 2, 3, 4]  # change as needed\nfor idx in sample_indices:\n    plot_final_sample(\n        wavelengths=wavelength,\n        y_true=y_final_true[idx],\n        y_pred=y_final_pred[idx],\n        sigma=sigma_final_calibrated[idx],\n        sample_id=idx\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:25:22.884622Z","iopub.execute_input":"2025-08-26T19:25:22.884917Z","iopub.status.idle":"2025-08-26T19:25:24.003205Z","shell.execute_reply.started":"2025-08-26T19:25:22.884894Z","shell.execute_reply":"2025-08-26T19:25:24.00256Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: build submission.csv (ALIGNED: FGS1 + 282 AIRS)\n# =========================\nimport os, time\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\n\n# ---- Base path & device ----\nBASE = os.getenv(\"BASE_URL\", \"/kaggle/input/ariel-data-challenge-2025\")\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", device)\n\n# ---- Load scalers (same as training) ----\nsc = np.load(\"models/scalers_v1.npz\")\nwc_min = float(sc[\"wc_min\"]); wc_max = float(sc[\"wc_max\"])\nymin   = float(sc[\"y_min\"]);  ymax   = float(sc[\"y_max\"])\ntargets_abs_max = float(sc[\"targets_abs_max\"])\ndata_abs_max    = float(sc[\"data_abs_max\"])\nprint(\"Scalers loaded.\")\n\ndef scale_wc_batch(x):        # x: (B,1,187)\n    return (x - wc_min) / (wc_max - wc_min + 1e-8)\n\ndef unscale_y_batch(y):       # y: (B,1) or (B,)\n    return y * (ymax - ymin + 1e-8) + ymin\n\n# ---- Load best weights ----\nnet1d = Net1D(dropout_rate=0.2).to(device)\nnet1d.load_state_dict(torch.load(\"models/net1d.pt\", map_location=device))\nnet1d.eval()\n\nnet2d = Net2D(time_len=40, wave_len=283, dropout=0.2).to(device)\nnet2d.load_state_dict(torch.load(\"models/net2d.pt\", map_location=device))\nnet2d.eval()\n\ndef enable_dropout_only(model: torch.nn.Module)->None:\n    \"\"\"Enable MC Dropout while keeping BatchNorm frozen.\"\"\"\n    for m in model.modules():\n        if isinstance(m, (torch.nn.Dropout, torch.nn.AlphaDropout)):\n            m.train()\n        elif isinstance(m, (torch.nn.BatchNorm1d, torch.nn.BatchNorm2d)):\n            m.eval()\n\n# ---- Test loader ----\nds_test  = ArielCNNDataset(BASE, split=\"test\", return_labels=False, disk_cache_dir=None)\ndl_test  = DataLoader(ds_test, batch_size=32, shuffle=False, num_workers=0,\n                      pin_memory=(device.type==\"cuda\"), persistent_workers=False)\nprint(\"Test size:\", len(ds_test))\n\n# ---- MC-Dropout inference ----\nS1D, S2D = 200, 50\nenable_dropout_only(net1d)\nenable_dropout_only(net2d)\n\nall_pid = []\nmu1d_list, sig1d_list = [], []\nmu2d_list, sig2d_list = [], []\n\nt0 = time.time()\nfor batch in dl_test:\n    pids = batch[\"pid\"]\n    all_pid.extend(pids)\n\n    # 1D: μ̂ and σ̂_μ   (FGS1 surrogate)\n    x1 = scale_wc_batch(batch[\"wc\"].to(device))          # (B,1,187)\n    samples_1d = []\n    for _ in range(S1D):\n        yhat = net1d(x1)                                 # (B,1) in [0,1]\n        samples_1d.append(yhat.detach().cpu().numpy())   # (B,1)\n    s1 = np.stack(samples_1d, axis=0)                    # (S,B,1)\n    mu1 = s1.mean(axis=0).squeeze(1)                     # (B,)\n    sg1 = s1.std(axis=0).squeeze(1)                      # (B,)\n    mu1 = unscale_y_batch(mu1)                           # (B,) original scale\n    sg1 = sg1 * (ymax - ymin + 1e-8)                     # (B,)\n\n    mu1d_list.append(mu1)\n    sig1d_list.append(sg1)\n\n    # 2D: Δ̂ and σ̂_Δ   (AIRS shape, 283 pix = 39..321)\n    x2 = (batch[\"map2d\"] / (data_abs_max + 1e-8)).to(device)  # (B,1,40,283)\n    samples_2d = []\n    for _ in range(S2D):\n        yhat2 = net2d(x2)                                  # (B,283) in [-1,1]\n        samples_2d.append(yhat2.detach().cpu().numpy())\n    s2 = np.stack(samples_2d, axis=0)                      # (S,B,283)\n    mu2 = s2.mean(axis=0) * (targets_abs_max + 1e-8)       # (B,283) -> original\n    sg2 = s2.std(axis=0)  * (targets_abs_max + 1e-8)       # (B,283)\n\n    mu2d_list.append(mu2)\n    sig2d_list.append(sg2)\n\nprint(f\"MC inference done in {time.time()-t0:.1f}s\")\n\n# ---- Concatenate over all batches ----\nmu_1d = np.concatenate(mu1d_list, axis=0)            # (N,)\nsig_1d = np.concatenate(sig1d_list, axis=0)          # (N,)\nmu_2d  = np.concatenate(mu2d_list, axis=0)           # (N,283) AIRS 39..321\nsig_2d = np.concatenate(sig2d_list, axis=0)          # (N,283)\n\n# ---- ALIGNED assembly for Kaggle ----\n# Kaggle expects 283 targets: [FGS1] + 282 AIRS points.\n# We KEEP AIRS 39..320 (drop the last pixel 321 -> tail drop).\nmu_2d_282  = mu_2d[:, :282]                          # (N,282) AIRS 39..320\nsig_2d_282 = sig_2d[:, :282]                         # (N,282)\n\n# wl_1 (FGS1) = μ̂_1d\n# wl_2..wl_283 = μ̂_1d + Δ̂ (for 282 AIRS wavelengths)\ny_pred_full = np.column_stack([mu_1d, mu_1d[:,None] + mu_2d_282])   # (N,283)\n\n# σ_1 (FGS) = σ̂_μ ;  σ_AIRS = sqrt(σ̂_μ^2 + σ̂_Δ^2)\nsigma_full  = np.column_stack([sig_1d, np.sqrt(sig_1d[:,None]**2 + sig_2d_282**2)])  # (N,283)\n\n# ---- Per-channel σ floor → global scale c → clip  ----\n# 1) compute per-column std from train.csv (wl_1..wl_283) and take 10% as a floor\ntrain_df = pd.read_csv(os.path.join(BASE, \"train.csv\"))\ntrain_specs = train_df[[f\"wl_{i}\" for i in range(1, 284)]].to_numpy(dtype=np.float64)  # (Ntrain, 283)\nper_col_std = train_specs.std(axis=0, ddof=1)                                          # (283,)\nsigma_floor = np.clip(0.1 * per_col_std, 8e-4, None)                                   # floor >= 8e-4\n\n# 2) global σ multiplier (RE-TUNE THIS AFTER ALIGNMENT; ~1.2–1.5 worked in your sweep)\nc_aligned = float(globals().get(\"calibration_factor_c_aligned\", 1.3))\n\nsigma_full = np.maximum(sigma_full, sigma_floor[None, :]) * c_aligned\nsigma_full = np.clip(sigma_full, 8e-4, 1e-2)  # final safety clip\n\n# ---- Build DataFrame in exact Kaggle order ----\nwl_cols = [f\"wl_{i}\"    for i in range(1, 284)]   # wl_1 = FGS1; wl_2..wl_283 = AIRS (39..320)\nsg_cols = [f\"sigma_{i}\" for i in range(1, 284)]\n\ndf_wl = pd.DataFrame(y_pred_full, columns=wl_cols)\ndf_sg = pd.DataFrame(sigma_full, columns=sg_cols)\ndf    = pd.concat([df_wl, df_sg], axis=1)\ndf.insert(0, \"planet_id\", all_pid)\n\n# ---- Strict checks ----\nassert df.shape[1] == 567, f\"Expect 567 columns, got {df.shape[1]}\"\nassert list(df.columns[:1]) == [\"planet_id\"]\nassert df.columns[1:284].tolist() == wl_cols\nassert df.columns[284:].tolist()   == sg_cols\nvals = df.iloc[:, 1:].to_numpy()\nassert np.isfinite(vals).all(), \"Found NaN/Inf\"\nsig_vals = df.filter(like=\"sigma_\").to_numpy()\nassert (sig_vals > 0).all(), \"All σ must be positive\"\nassert df[\"planet_id\"].nunique() == len(df), \"Duplicate planet_id rows\"\n\n# ---- Write file ----\nout_path = \"/kaggle/working/submission.csv\"  # or your local path\ndf.to_csv(out_path, index=False)\nprint(\"Saved:\", out_path, \"shape:\", df.shape)\nprint(\"NOTE: In the Kaggle submit dialog, select 'submission.csv' from Output files.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-26T19:25:24.004008Z","iopub.execute_input":"2025-08-26T19:25:24.004249Z","iopub.status.idle":"2025-08-26T19:25:27.239409Z","shell.execute_reply.started":"2025-08-26T19:25:24.004233Z","shell.execute_reply":"2025-08-26T19:25:27.238757Z"}},"outputs":[],"execution_count":null}]}