{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":12934072,"sourceType":"datasetVersion","datasetId":8184689},{"sourceId":555288,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":422712,"modelId":440272},{"sourceId":555724,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":422840,"modelId":440388},{"sourceId":556118,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":423136,"modelId":440674},{"sourceId":562345,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":425483,"modelId":442977},{"sourceId":562382,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":425511,"modelId":443004},{"sourceId":562801,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":425850,"modelId":443339},{"sourceId":565290,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":426839,"modelId":443877},{"sourceId":566246,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":427060,"modelId":444079},{"sourceId":569013,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":427689,"modelId":444695},{"sourceId":571345,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":429220,"modelId":446197},{"sourceId":572460,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":429428,"modelId":446389},{"sourceId":576673,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":431189,"modelId":448112},{"sourceId":576707,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":431200,"modelId":448125},{"sourceId":578530,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":431930,"modelId":448846},{"sourceId":580225,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":433013,"modelId":449895},{"sourceId":580229,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":433017,"modelId":449899},{"sourceId":580323,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":433099,"modelId":449979},{"sourceId":582357,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":434802,"modelId":451653},{"sourceId":582359,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":434804,"modelId":451655},{"sourceId":582361,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":434806,"modelId":451657},{"sourceId":583446,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":435747,"modelId":452548},{"sourceId":583453,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":435754,"modelId":452555},{"sourceId":583455,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":435756,"modelId":452557},{"sourceId":588780,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":440253,"modelId":456799},{"sourceId":589669,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":440969,"modelId":457512}],"dockerImageVersionId":31090,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"***Introduction*** \n\nThis notebook is focused on the test data pipeline and submission formatting for the Ariel Data Challenge. For each testing file, we perform all necessary preprocessing steps, including detector calibration, bad pixel masking, median filtering, inpainting, and standardized binning, before running model inference and assembling the results into the final submission file.\n\nPlanets with multiple observations are processed separately. We also use a stride of 10 to process each signal 5 times (all of the signals have been median filtered with kernel size 101 to remove outliers but retain necessary details). To ensemble we run the data through as many models as we wish to. The final, best performing results were from a simple average of both the values and the sigmas, though other ensemble methods were explored.\n\nPreprocessing routines draw on techniques and code from the [binning and processing notebook](https://www.kaggle.com/code/gordonyip/update-calibrating-and-binning-astronomical-data). Changes were made to increase the efficiency, change the masking behavior, and implement inpainting. We also rely on the ruptures Python library for change point detection.","metadata":{}},{"cell_type":"code","source":"#!pip install ruptures","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\n#loading_time = time.perf_counter() - start_loading\n#print(f\"Data loading timeC: {loading_time:.4f} sec\")\n#start_loading = time.perf_counter()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:48:05.453545Z","iopub.execute_input":"2025-09-24T23:48:05.454018Z","iopub.status.idle":"2025-09-24T23:48:05.466583Z","shell.execute_reply.started":"2025-09-24T23:48:05.453994Z","shell.execute_reply":"2025-09-24T23:48:05.465912Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from: https://www.kaggle.com/code/gordonyip/update-calibrating-and-binning-astronomical-data\nimport numpy as np\nimport pandas as pd\nimport itertools\nimport os\nimport glob \nfrom astropy.stats import sigma_clip\nfrom tqdm import tqdm\nimport re\nfrom skimage.restoration import inpaint_biharmonic\nimport torch\nimport matplotlib.pyplot as plt\nfrom scipy.signal import medfilt\n#import ruptures as rpt\nfrom scipy.signal import savgol_filter","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:48:08.088613Z","iopub.execute_input":"2025-09-24T23:48:08.089428Z","iopub.status.idle":"2025-09-24T23:48:08.093664Z","shell.execute_reply.started":"2025-09-24T23:48:08.089401Z","shell.execute_reply":"2025-09-24T23:48:08.092843Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def ADC_convert(signal, gain, offset):\n    signal = signal.astype(np.float64)\n    signal /= gain\n    signal += offset\n    return signal\n\ndef mask_hot_dead_slow(signal, dead, dark):\n    hot = sigma_clip(\n        dark, sigma=5, maxiters=5\n    ).mask\n    hot = np.tile(hot, (signal.shape[0], 1, 1))\n    dead = np.tile(dead, (signal.shape[0], 1, 1))\n    signal = np.ma.masked_where(dead, signal)\n    signal = np.ma.masked_where(hot, signal)\n    return signal\n\nimport numpy as np\nfrom astropy.stats import sigma_clip\n\ndef mask_hot_dead(signal, dead, dark):\n    # Compute hot mask only once, cache if multiple calls with same dark\n    hot_mask_2d = sigma_clip(dark, sigma=5, maxiters=5).mask\n    \n    # Broadcast masks without np.tile to save memory/time\n    # shape (time, x, y)\n    hot_mask = np.broadcast_to(hot_mask_2d, signal.shape)\n    dead_mask = np.broadcast_to(dead, signal.shape)\n    \n    combined_mask = np.logical_or(hot_mask, dead_mask)\n    \n    # Construct new masked array in one step\n    masked_signal = np.ma.array(signal, mask=combined_mask)\n    \n    return masked_signal\n\ndef apply_linear_corr(linear_corr,clean_signal):\n    linear_corr = np.flip(linear_corr, axis=0)\n    for x, y in itertools.product(\n                range(clean_signal.shape[1]), range(clean_signal.shape[2])\n            ):\n        poli = np.poly1d(linear_corr[:, x, y])\n        clean_signal[:, x, y] = poli(clean_signal[:, x, y])\n    return clean_signal\n\ndef clean_dark_old(signal, dead, dark, dt):\n\n    dark = np.ma.masked_where(dead, dark)\n    dark = np.tile(dark, (signal.shape[0], 1, 1))\n\n    signal -= dark* dt[:, np.newaxis, np.newaxis]\n    return signal\n\ndef clean_dark(signal, dead, dark, dt):\n    masked_dark = np.ma.masked_array(dark, mask=dead)  # mask dark where dead\n    # Broadcast dark along time axis and multiply by dt, then subtract\n    signal -= masked_dark * dt[:, np.newaxis, np.newaxis]\n    return signal\n\ndef get_cds(signal):\n    cds = signal[:,1::2,:,:] - signal[:,::2,:,:]\n    return cds\n\ndef correct_flat_field_old(flat,dead, signal):\n    flat = flat.transpose(1, 0)\n    dead = dead.transpose(1, 0)\n    flat = np.ma.masked_where(dead, flat)\n    flat = np.tile(flat, (signal.shape[0], 1, 1))\n    signal = signal / flat\n    return signal\n    \ndef correct_flat_field(flat,dead, signal):\n    flat = flat.transpose(1, 0)  # shape (cols, rows)\n    dead = dead.transpose(1, 0)\n    flat = np.ma.masked_where(dead, flat)\n    flat = flat[np.newaxis, :, :]  # add axis for broadcast\n    signal = signal / flat  # broadcasting instead of tile\n    return signal","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:48:09.703781Z","iopub.execute_input":"2025-09-24T23:48:09.704275Z","iopub.status.idle":"2025-09-24T23:48:09.714275Z","shell.execute_reply.started":"2025-09-24T23:48:09.704247Z","shell.execute_reply":"2025-09-24T23:48:09.713604Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from ruptures https://centre-borelli.github.io/ruptures-docs/code-reference/detection/binseg-reference/#ruptures.detection.binseg.Binseg\nimport abc\nfrom functools import lru_cache\n\nfrom itertools import tee\n\ndef pairwise(iterable):\n    \"s -> (s0,s1), (s1,s2), (s2, s3), ...\"\n    a, b = tee(iterable)\n    next(b, None)\n    return zip(a, b)\n\ndef sanity_check(n_samples, n_bkps, jump, min_size):\n    \"\"\"Check that segmentation parameters are valid.\n\n    Args:\n        n_samples (int): number of samples in the signal\n        n_bkps (int): number of requested breakpoints\n        jump (int): subsample jump size\n        min_size (int): minimum segment size\n\n    Returns:\n        bool: True if parameters are consistent; False otherwise\n    \"\"\"\n    if n_samples < 1:\n        return False\n    if n_bkps < 0:\n        return False\n    if jump < 1:\n        return False\n    if min_size < 1:\n        return False\n    # Check that at least one segment can be formed\n    if n_samples < (n_bkps + 1) * min_size:\n        return False\n    return True\n\n\nclass BaseCost(object, metaclass=abc.ABCMeta):\n    \"\"\"Base class for all segment cost classes.\n\n    Notes:\n        All classes should specify all the parameters that can be set\n        at the class level in their ``__init__`` as explicit keyword\n        arguments (no ``*args`` or ``**kwargs``).\n    \"\"\"\n\n    @abc.abstractmethod\n    def fit(self, *args, **kwargs):\n        \"\"\"Set the parameters of the cost function, for instance the Gram\n        matrix, etc.\"\"\"\n        pass\n\n    @abc.abstractmethod\n    def error(self, start, end):\n        \"\"\"Returns the cost on segment [start:end].\"\"\"\n        pass\n\n    def sum_of_costs(self, bkps):\n        \"\"\"Returns the sum of segments cost for the given segmentation.\n\n        Args:\n            bkps (list): list of change points. By convention, bkps[-1]==n_samples.\n\n        Returns:\n            float: sum of costs\n        \"\"\"\n        soc = sum(self.error(start, end) for start, end in pairwise([0] + bkps))\n        return soc\n\n    @property\n    @abc.abstractmethod\n    def model(self):\n        pass\n\ndef cost_factory(model=\"l2\", **params):\n    # Example placeholder. You must implement cost classes (like CostL2) yourself!\n    if model == \"l2\":\n        return CostL2(**params)\n    elif model == \"l1\":\n        return CostL1(**params)\n    # Add other models as needed\n    else:\n        raise ValueError(f\"Unknown model '{model}'\")\n\nclass CostL2(BaseCost):\n    r\"\"\"Least squared deviation.\"\"\"\n\n    model = \"l2\"\n\n    def __init__(self):\n        \"\"\"Initialize the object.\"\"\"\n        self.signal = None\n        self.min_size = 1\n\n    def fit(self, signal) -> \"CostL2\":\n        \"\"\"Set parameters of the instance.\n\n        Args:\n            signal (array): array of shape (n_samples,) or (n_samples, n_features)\n\n        Returns:\n            self\n        \"\"\"\n        if signal.ndim == 1:\n            self.signal = signal.reshape(-1, 1)\n        else:\n            self.signal = signal\n\n        return self\n\n    def error(self, start, end) -> float:\n        \"\"\"Return the approximation cost on the segment [start:end].\n\n        Args:\n            start (int): start of the segment\n            end (int): end of the segment\n\n        Returns:\n            segment cost\n\n        Raises:\n            NotEnoughPoints: when the segment is too short (less than `min_size` samples).\n        \"\"\"\n        if end - start < self.min_size:\n            raise NotEnoughPoints\n\n        return self.signal[start:end].var(axis=0).sum() * (end - start)\n\nclass BaseEstimator(metaclass=abc.ABCMeta):\n    \"\"\"Base class for all change point detection estimators.\n\n    Notes:\n        All estimators should specify all the parameters that can be set\n        at the class level in their ``__init__`` as explicit keyword\n        arguments (no ``*args`` or ``**kwargs``).\n    \"\"\"\n\n    @abc.abstractmethod\n    def fit(self, *args, **kwargs):\n        \"\"\"To call the segmentation algorithm.\"\"\"\n        pass\n\n    @abc.abstractmethod\n    def predict(self, *args, **kwargs):\n        \"\"\"To call the segmentation algorithm.\"\"\"\n        pass\n\n    @abc.abstractmethod\n    def fit_predict(self, *args, **kwargs):\n        \"\"\"To call the segmentation algorithm.\"\"\"\n        pass\n\nclass Binseg(BaseEstimator):\n    \"\"\"Binary segmentation.\"\"\"\n\n    def __init__(self, model=\"l2\", custom_cost=None, min_size=2, jump=5, params=None):\n        \"\"\"Initialize a Binseg instance.\n\n        Args:\n            model (str, optional): segment model, [\"l1\", \"l2\", \"rbf\",...]. Not used if ``'custom_cost'`` is not None.\n            custom_cost (BaseCost, optional): custom cost function. Defaults to None.\n            min_size (int, optional): minimum segment length. Defaults to 2 samples.\n            jump (int, optional): subsample (one every *jump* points). Defaults to 5 samples.\n            params (dict, optional): a dictionary of parameters for the cost instance.\n        \"\"\"\n        if custom_cost is not None and isinstance(custom_cost, BaseCost):\n            self.cost = custom_cost\n        else:\n            if params is None:\n                self.cost = cost_factory(model=model)\n            else:\n                self.cost = cost_factory(model=model, **params)\n        self.min_size = max(min_size, self.cost.min_size)\n        self.jump = jump\n        self.n_samples = None\n        self.signal = None\n\n    def _seg(self, n_bkps=None, pen=None, epsilon=None):\n        \"\"\"Computes the binary segmentation.\n\n        The stopping rule depends on the parameter passed to the function.\n\n        Args:\n            n_bkps (int): number of breakpoints to find before stopping.\n            penalty (float): penalty value (>0)\n            epsilon (float): reconstruction budget (>0)\n\n        Returns:\n            dict: partition dict {(start, end): cost value,...}\n        \"\"\"\n        # initialization\n        bkps = [self.n_samples]\n        stop = False\n        while not stop:\n            stop = True\n            new_bkps = [\n                self.single_bkp(start, end) for start, end in pairwise([0] + bkps)\n            ]\n            bkp, gain = max(new_bkps, key=lambda x: x[1])\n\n            if bkp is None:  # all possible configuration have been explored.\n                break\n\n            if n_bkps is not None:\n                if len(bkps) - 1 < n_bkps:\n                    stop = False\n            elif pen is not None:\n                if gain > pen:\n                    stop = False\n            elif epsilon is not None:\n                error = self.cost.sum_of_costs(bkps)\n                if error > epsilon:\n                    stop = False\n\n            if not stop:\n                bkps.append(bkp)\n                bkps.sort()\n        partition = {\n            (start, end): self.cost.error(start, end)\n            for start, end in pairwise([0] + bkps)\n        }\n        return partition\n\n    @lru_cache(maxsize=None)\n    def single_bkp(self, start, end):\n        \"\"\"Return the optimal breakpoint of [start:end] (if it exists).\"\"\"\n        segment_cost = self.cost.error(start, end)\n        if np.isinf(segment_cost) and segment_cost < 0:  # if cost is -inf\n            return None, 0\n        gain_list = list()\n        for bkp in range(start, end, self.jump):\n            if bkp - start >= self.min_size and end - bkp >= self.min_size:\n                gain = (\n                    segment_cost\n                    - self.cost.error(start, bkp)\n                    - self.cost.error(bkp, end)\n                )\n                gain_list.append((gain, bkp))\n        try:\n            gain, bkp = max(gain_list)\n        except ValueError:  # if empty sub_sampling\n            return None, 0\n        return bkp, gain\n\n    def fit(self, signal) -> \"Binseg\":\n        \"\"\"Compute params to segment signal.\n\n        Args:\n            signal (array): signal to segment. Shape (n_samples, n_features) or (n_samples,).\n\n        Returns:\n            self\n        \"\"\"\n        # update some params\n        if signal.ndim == 1:\n            self.signal = signal.reshape(-1, 1)\n        else:\n            self.signal = signal\n        self.n_samples, _ = self.signal.shape\n        self.cost.fit(signal)\n        self.single_bkp.cache_clear()\n\n        return self\n\n    def predict(self, n_bkps=None, pen=None, epsilon=None):\n        \"\"\"Return the optimal breakpoints.\n\n        Must be called after the fit method. The breakpoints are associated with the\n        signal passed to [`fit()`][ruptures.detection.binseg.Binseg.fit].\n        The stopping rule depends on the parameter passed to the function.\n\n        Args:\n            n_bkps (int): number of breakpoints to find before stopping.\n            pen (float): penalty value (>0)\n            epsilon (float): reconstruction budget (>0)\n\n        Raises:\n            AssertionError: if none of `n_bkps`, `pen`, `epsilon` is set.\n            BadSegmentationParameters: in case of impossible segmentation\n                configuration\n\n        Returns:\n            list: sorted list of breakpoints\n        \"\"\"\n        msg = \"Give a parameter.\"\n        assert any(param is not None for param in (n_bkps, pen, epsilon)), msg\n\n        # raise an exception in case of impossible segmentation configuration\n        if not sanity_check(\n            n_samples=self.cost.signal.shape[0],\n            n_bkps=0 if n_bkps is None else n_bkps,\n            jump=self.jump,\n            min_size=self.min_size,\n        ):\n            raise BadSegmentationParameters\n\n        partition = self._seg(n_bkps=n_bkps, pen=pen, epsilon=epsilon)\n        bkps = sorted(e for s, e in partition.keys())\n        return bkps\n\n    def fit_predict(self, signal, n_bkps=None, pen=None, epsilon=None):\n        \"\"\"Fit to the signal and return the optimal breakpoints.\n\n        Helper method to call fit and predict once\n\n        Args:\n            signal (array): signal. Shape (n_samples, n_features) or (n_samples,).\n            n_bkps (int): number of breakpoints.\n            pen (float): penalty value (>0)\n            epsilon (float): reconstruction budget (>0)\n\n        Returns:\n            list: sorted list of breakpoints\n        \"\"\"\n        self.fit(signal)\n        return self.predict(n_bkps=n_bkps, pen=pen, epsilon=epsilon)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:48:11.247763Z","iopub.execute_input":"2025-09-24T23:48:11.248016Z","iopub.status.idle":"2025-09-24T23:48:11.268827Z","shell.execute_reply.started":"2025-09-24T23:48:11.247998Z","shell.execute_reply":"2025-09-24T23:48:11.26804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## we will start by getting the index of the training data:\ndef get_index(files,CHUNKS_SIZE ):\n    index = []\n    for file in files :\n        file_name = file.split('/')[-1]\n        if file_name.split('_')[0] == 'AIRS-CH0' and file_name.split('_')[-1] == '0.parquet':\n            file_index = os.path.basename(os.path.dirname(file))\n            index.append(int(file_index))\n    index = np.array(index)\n    index = np.sort(index) \n    # credit to DennisSakva\n    index=np.array_split(index, len(index)//CHUNKS_SIZE)\n    \n    return index\n\ndef get_multiobs_index(files, CHUNKS_SIZE):\n    \"\"\"\n    Extract (planet_id, obs_num) pairs from AIRS-CH0_signal_X.parquet files.\n    Returns: list of (planet_id, obs_num) tuples in sorted order, split into chunks.\n    \"\"\"\n    index = []\n    # Regex: AIRS-CH0_signal_{obs}.parquet\n    pattern = re.compile(r'^AIRS-CH0_signal_(\\d+)\\.parquet$')\n    for file in files:\n        file_name = os.path.basename(file)\n        match = pattern.match(file_name)\n        if match:\n            planet_id = os.path.basename(os.path.dirname(file))\n            obs_num = int(match.group(1))\n            index.append((int(planet_id), obs_num))\n    # Optional: sort by planet then obs number\n    index.sort()\n    # Remove duplicates in case of any\n    index = list(dict.fromkeys(index))\n    if len(index) >= CHUNKS_SIZE and CHUNKS_SIZE > 0:\n        index_chunks = np.array_split(index, len(index)//CHUNKS_SIZE)\n    else:\n        index_chunks = [index]\n    return index_chunks\n\ndef bin_obs(arr, binning, axis=1):\n    # Ensure input is a masked array\n    bin_size = binning\n    arr = np.ma.masked_array(arr)\n    shape = list(arr.shape)\n    n_bins = shape[axis] // bin_size\n    new_shape = shape[:axis] + [n_bins, bin_size] + shape[axis+1:]\n    arr_reshaped = np.ma.reshape(arr, new_shape)\n    # Now sum along the bin_size axis, which is axis=axis+1\n    return np.ma.sum(arr_reshaped, axis=axis+1)\n\ndef median_filter_time(masked_arr, kernel_size=3):\n    \"\"\"Apply 1D median filter (default: size 3) along time axis (axis=1) for each batch.\n    Ignores masked voxels; uses available neighbors at edges. Preserves masked array structure.\"\"\"\n    assert kernel_size % 2 == 1, \"Kernel size must be odd!\"\n    batch_dim, time_dim, X, Y = masked_arr.shape\n    pad = kernel_size // 2\n    result = np.ma.masked_all(masked_arr.shape, dtype=masked_arr.dtype)\n    arr_data = masked_arr.data\n    arr_mask = masked_arr.mask\n\n    for b in range(batch_dim):\n        for t in range(time_dim):\n            lo = max(0, t - pad)\n            hi = min(time_dim, t + pad + 1)\n            window = arr_data[b, lo:hi, :, :]\n            window_mask = arr_mask[b, lo:hi, :, :]\n            window_ma = np.ma.masked_array(window, mask=window_mask)\n            # Use np.ma.median as a function for better compatibility\n            median_vals = np.ma.median(window_ma, axis=0)\n            result.data[b, t] = median_vals.data\n            result.mask[b, t] = median_vals.mask\n\n    return result\n\ndef already_saved(chunk_name, path_out):\n    airs_file = os.path.join(path_out, f'AIRS_clean_train_{chunk_name}.pt')\n    fgs1_file = os.path.join(path_out, f'FGS1_clean_train_{chunk_name}.pt')\n    return os.path.exists(airs_file) and os.path.exists(fgs1_file)\n\ndef median_filter_and_downsample(\n    signal,\n    median_filter_window=101,\n    stride=10,\n    title='Median Filtered and Downsampled Signal',\n    plot = True\n):\n    \"\"\"\n    Applies a median filter to a 1D signal, crops edges, downsamples by specified stride, \n    and plots the result. Returns the downsampled signal and its x coordinates.\n    \"\"\"\n    # Apply median filter (pads internally)\n    window_size = median_filter_window  # must be odd\n    border = (window_size - 1) // 2\n\n    median_filtered_full = medfilt(signal, kernel_size=window_size)\n\n    # Crop edges to remove padding artifacts\n    median_filtered_cropped = median_filtered_full[border:-border]\n    x_cropped = np.arange(border, len(signal) - border)\n\n    # Downsample\n    downsampled_signal = median_filtered_cropped[::stride]\n    x_downsampled = x_cropped[::stride]\n\n    if plot:\n        # Plot\n        plt.figure(figsize=(14, 7))\n        plt.plot(x_cropped, median_filtered_cropped, label=f'Median Filtered (window={window_size}, stride=1)', linewidth=2)\n        plt.plot(x_downsampled, downsampled_signal, marker='o', linestyle='--',\n                 label=f'Filtered & Downsampled (stride={stride})')\n        plt.title(title)\n        plt.xlabel('Sample Index')\n        plt.ylabel('Signal Value')\n        plt.legend()\n        plt.grid(True)\n        plt.tight_layout()\n        plt.show()\n\n    return downsampled_signal, x_downsampled\n\ndef plot_transit_edges(\n    signal,\n    window_length=15,\n    polyorder=2,\n    window=5,\n    percentile=30,\n    min_size=10,\n    title=\"Local Transit Edges Detection (Split at Min)\",\n    plot_raw=True,\n    plot = True\n):\n    \"\"\"\n    Plots transit edges and detected change points for a 1D signal.\n    Returns onset index (left edge), offset index (right edge).\n    \"\"\"\n    # Smoothing\n    smoothed_signal = savgol_filter(signal, window_length=window_length, polyorder=polyorder)\n    #smoothed_signal = signal\n    \n    # Find global minimum (likely transit midpoint)\n    min_index = np.argmin(smoothed_signal)\n\n    # Split signal at minimum\n    signal_left = smoothed_signal[:min_index]\n    signal_right = smoothed_signal[min_index:]\n\n    # Detect on left half (before transit: drop)\n    n_bkps = 1\n    algo_left = Binseg(model=\"l2\", min_size=min_size).fit(signal_left)\n    bkps_left = algo_left.predict(n_bkps=n_bkps)\n    change_left = bkps_left[0]\n\n    onset_left = find_transit_edge_local(signal_left, change_left, find_onset=True, window=window, percentile=percentile)\n\n    # Detect on right half (after transit: rise)\n    algo_right = Binseg(model=\"l2\", min_size=min_size).fit(signal_right)\n    bkps_right = algo_right.predict(n_bkps=n_bkps)\n    change_right = bkps_right[0]\n\n    offset_right = find_transit_edge_local(signal_right, change_right, find_onset=False, window=window, percentile=percentile)\n    offset_right_global = min_index + offset_right\n\n    # (Optional) change points for info\n    midpoints = [change_left, min_index + change_right]\n\n    # Plot for confirmation\n    if plot:\n        plt.figure(figsize=(12, 6))\n        if plot_raw:\n            plt.plot(signal, label='Raw signal', color='gray', alpha=0.4)\n        plt.plot(smoothed_signal, label='Smoothed signal', color='navy')\n        plt.axvline(min_index, color='black', linestyle='--', label='Transit Min')\n\n        plt.axvline(onset_left, color='green', linestyle='-', label='Onset (start drop)', lw=3)\n        plt.scatter([onset_left], smoothed_signal[[onset_left]], color='green', s=80, zorder=10)\n\n        plt.axvline(offset_right_global, color='red', linestyle='-', label='Offset (end rise)', lw=3)\n        plt.scatter([offset_right_global], smoothed_signal[[offset_right_global]], color='red', s=80, zorder=10)\n\n        plt.axvline(midpoints[0], color='purple', linestyle='--', label='Change point (start)')\n        plt.axvline(midpoints[1], color='purple', linestyle='--', label='Change point (end)')\n\n        plt.legend()\n        plt.xlabel('Sample Index')\n        plt.ylabel('Signal Value')\n        plt.title(title)\n        plt.tight_layout()\n        plt.show()\n\n    #print(f\"Onset index (left side): {onset_left}\")\n    #print(f\"Offset index (right side, global): {offset_right_global}\")\n    return onset_left, offset_right_global, min_index, np.min(smoothed_signal)\n\ndef find_transit_edge_local(signal, change_point, find_onset=True, window=5, percentile=80):\n    if find_onset:\n        region = signal[:change_point]\n        threshold = np.percentile(region, percentile)\n        for i in range(change_point, window, -1):\n            if np.all(signal[i-window:i] >= threshold):\n                return i\n        return window\n    else:\n        region = signal[change_point:]\n        threshold = np.percentile(region, percentile)\n        for i in range(change_point, len(signal)-window):\n            if np.all(signal[i:i+window] >= threshold):\n                return i\n        return len(signal)-window\n\ndef fit_and_plot_baseline(\n    signal,\n    onset_idx,\n    offset_idx,\n    delta=0,\n    degree=2,\n    planet_id=None,\n    plot=True,\n    title='Baseline Fit'\n):\n    \"\"\"\n    Fit a polynomial baseline curve to regions outside [onset_idx, offset_idx],\n    with delta applied to edges. Plots result optionally.\n    Returns: fitted_curve, coeffs, idx_baseline\n    \"\"\"\n    # Adjust edges with delta\n    phase1 = max(0, onset_idx - delta)\n    phase2 = min(len(signal), offset_idx + delta)\n    \n    # Indices for left and right baseline regions\n    idx_left = np.arange(0, phase1)\n    idx_right = np.arange(phase2, len(signal))\n    idx_baseline = np.concatenate([idx_left, idx_right])\n    y_baseline = signal[idx_baseline]\n    \n    # Get a boolean mask for valid values (not NaN, not Inf)\n    valid_mask = (~np.isnan(y_baseline)) & (~np.isinf(y_baseline))\n\n    # Filter both arrays\n    idx_baseline = idx_baseline[valid_mask]\n    y_baseline = y_baseline[valid_mask]\n    \n    # Fit polynomial\n    coeffs = np.polyfit(idx_baseline, y_baseline, deg=degree)\n    poly = np.poly1d(coeffs)\n    fitted_curve = poly(np.arange(len(signal)))\n    \n    # Plotting\n    if plot:\n        plt.figure(figsize=(10, 4))\n        plt.plot(signal, label='Signal')\n        plt.plot(fitted_curve, '--', label='Fitted Baseline', color='orange')\n        plt.scatter(idx_baseline, signal[idx_baseline], color='green', label='Baseline Points')\n        plt.axvline(onset_idx, color='r', linestyle='--', label='Transit Onset', lw=2)\n        plt.axvline(offset_idx, color='b', linestyle='--', label='Transit Offset', lw=2)\n        plt.axvspan(phase1, min(len(signal), onset_idx + delta), color='r', alpha=0.2, label='Delta Onset')\n        plt.axvspan(max(0, offset_idx - delta), phase2, color='b', alpha=0.2, label='Delta Offset')\n        plt.legend()\n        sub_id = f\" for Planet ID {planet_id}\" if planet_id is not None else \"\"\n        plt.title(f'{title}{sub_id}')\n        plt.xlabel('Sample Index')\n        plt.ylabel('Signal')\n        plt.tight_layout()\n        plt.show()\n    \n    return fitted_curve, coeffs, idx_baseline","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:48:14.986903Z","iopub.execute_input":"2025-09-24T23:48:14.987175Z","iopub.status.idle":"2025-09-24T23:48:15.012285Z","shell.execute_reply.started":"2025-09-24T23:48:14.987156Z","shell.execute_reply":"2025-09-24T23:48:15.011763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"path_folder = '/kaggle/input/ariel-data-challenge-2025' # path to the folder containing the data\npath_out = '/kaggle/working/processed_datak'\nos.makedirs(path_out, exist_ok=True)\nfiles = glob.glob(os.path.join(path_folder, 'test','*','*'))\n\nCHUNKS_SIZE = 1\nindex_chunks = get_multiobs_index(files, CHUNKS_SIZE)\n\ntrain_adc_info = pd.read_csv(os.path.join(path_folder, 'adc_info.csv'))\naxis_info = pd.read_parquet(os.path.join(path_folder,'axis_info.parquet'))\nDO_MASK = True\nDO_THE_NL_CORR = False\nDO_DARK = True\nDO_FLAT = True\nTIME_BINNING = True\nFILT = False\n\ncut_inf, cut_sup = 0, 356\nl = cut_sup - cut_inf\ncount = 0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:48:17.656572Z","iopub.execute_input":"2025-09-24T23:48:17.65689Z","iopub.status.idle":"2025-09-24T23:48:17.853305Z","shell.execute_reply.started":"2025-09-24T23:48:17.656871Z","shell.execute_reply":"2025-09-24T23:48:17.852694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(device)\n\nconstants = torch.load('/kaggle/input/norm05/normalization_constants05.pt')\nconstantMEANS = torch.tensor([constants[\"MRS\"],constants[\"MMS\"],constants['MTS'],constants['MMP'],constants['MP'],constants['MSMA'],constants['MI']]).to(device)\nconstantSTDS = torch.tensor([constants[\"SRS\"],constants[\"SMS\"],constants['STS'],constants['SMP'],constants['SP'],constants['SSMA'],constants['SI']]).to(device)\nconstantMMD = constants[\"MMD\"].to(device)\nconstantSMD = constants[\"SMD\"].to(device)\nconstantMSD = constants[\"MSD\"].to(device)\nconstantSSD = constants[\"SSD\"].to(device)\n\ndef get_planet_values(star_info_df, planet_id_list):\n    columns = ['Rs', 'Ms', 'Ts', 'Mp', 'P', 'sma', 'i']\n    filtered_df = star_info_df[star_info_df['planet_id'].isin(planet_id_list)]\n    values_list = filtered_df[columns].values.tolist()\n    return values_list\n\nclass MiniReducerSpatialBoth(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super().__init__()\n        # Collapse both spatial axes: (32,32) -> (1,32)\n        self.conv_spatial = nn.Conv3d(\n            in_channels, in_channels,\n            kernel_size=(1, 27, 27),\n            stride=(1, 1, 1),\n            padding=0\n        )\n        # Linear projection of flattened spatial patch to 32 features\n        self.lin = nn.Linear(36, 32)\n\n    def forward(self, x):\n        #print(f\"Input: {x.shape}\")                  # (B, C, T, 32, 32)\n        x = self.conv_spatial(x)\n        #print(f\"After spatial collapse: {x.shape}\") # (B, C, T, H', W')\n        x = F.relu(x)\n        B, C, T, H, W = x.shape\n        x = x.view(B, C, T, H * W)                  # Flatten spatial dims\n        x = x.reshape(-1, H * W)                    # Merge (B, C, T) for linear layer\n        x = self.lin(x)                             # Project to 32\n        x = x.view(B, C, T, 32)                     # Restore shape\n        x = x.unsqueeze(3)                          # (B, C, T, 1, 32)\n        #print(x.shape)\n        return x\n\nclass ResidualBlock3D(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size, pool_kernel, pool_stride,\n                 dilation=1, circular_pad_wavelength=False):\n        super().__init__()\n        # Allow tuple for dilation (pythonic for axis control)\n        if isinstance(dilation, int):\n            dilation = (dilation, dilation, dilation)\n        # Compute per-axis padding\n        if isinstance(kernel_size, int):\n            kernel_size = (kernel_size, kernel_size, kernel_size)\n        pad = tuple(d * (k // 2) for d, k in zip(dilation, kernel_size))  # (D, H, W)\n        self.pad = pad\n        self.circular_pad_wavelength = circular_pad_wavelength\n\n        self.conv = nn.Conv3d(\n            in_channels, out_channels,\n            kernel_size=kernel_size,\n            padding=0 if circular_pad_wavelength else pad,  # all-zero padding applied manually if circular\n            dilation=dilation\n        )\n        self.bn = nn.BatchNorm3d(out_channels)\n        self.pool = nn.MaxPool3d(kernel_size=pool_kernel, stride=pool_stride)\n        self.match_channels = None\n        if in_channels != out_channels:\n            self.match_channels = nn.Conv3d(in_channels, out_channels, kernel_size=1)\n        \n    def forward(self, x):\n        identity = x\n        out = x\n        # Apply circular padding ONLY to wavelength axis\n        if self.circular_pad_wavelength:\n            # self.pad = (pad_D, pad_H, pad_W)\n            pad_D, pad_H, pad_W = self.pad\n            # F.pad expects (W_left, W_right, H_top, H_bottom, D_front, D_back)\n            # Wavelength=H axis (axis=3)\n            out = F.pad(out, (0,0, pad_H, pad_H, 0,0), mode=\"circular\")\n            # add zero padding on depth and width if needed\n            if pad_D > 0 or pad_W > 0:\n                out = F.pad(out, (pad_W, pad_W, 0,0, pad_D, pad_D), mode=\"constant\", value=0)\n        out = self.conv(out)\n        out = self.bn(out)\n        out = F.relu(out)\n        out = self.pool(out)\n        identity_pooled = self.pool(identity)\n        if self.match_channels:\n            identity_pooled = self.match_channels(identity_pooled)\n        out = out + identity_pooled\n        return out\n\n\nclass Custom2DCNN(nn.Module):\n    def __init__(self, in_channels=2,\n                 conv_kernel_sizes=[3,3,3,3,3,3],\n                 conv_filters=[32,64,128,256,512,1028],\n                 pool_kernel_sizes=[[16,1,4],[16,1,4],[16,1,4],[16,1,4],[16,1,4],[16,1,4]],\n                 pool_strides=[[8,1,2],[8,1,2],[8,1,2],[8,1,2],[8,1,2],[8,1,2]],\n                 dilation_rates=None,\n                 circular_pad_wavelength_layers=None):\n        super().__init__()\n        self.blocks = nn.ModuleList()\n        chans = in_channels\n        if dilation_rates is None:\n            dilation_rates = [1] * len(conv_filters)\n        if circular_pad_wavelength_layers is None:\n            circular_pad_wavelength_layers = [False] * len(conv_filters)\n        for idx in range(len(conv_filters)):\n            self.blocks.append(\n                ResidualBlock3D(\n                    chans, conv_filters[idx],\n                    kernel_size=conv_kernel_sizes[idx],\n                    pool_kernel=pool_kernel_sizes[idx],\n                    pool_stride=pool_strides[idx],\n                    dilation=dilation_rates[idx],\n                    circular_pad_wavelength=circular_pad_wavelength_layers[idx]\n                )\n            )\n            chans = conv_filters[idx]\n        self.fc1 = nn.Linear(chans+357+357+7, 512)\n        self.fc_x = nn.Linear(512, 283)\n        self.fc_y = nn.Linear(512, 283)\n        self.fc_sigma = nn.Linear(512, 283)\n        #self.fgs1_reducer = MiniReducerSpatialBoth(in_channels=2, out_channels=2)\n    def forward(self, data, planet_data):\n        xmean = torch.mean(data,dim=3).unsqueeze(3)\n        xstd = torch.std(data,dim=3).unsqueeze(3)\n        x = (data-xmean)/xstd\n        xm = (xmean - constantMMD.view(1, 1, 357, 1))/constantSMD.view(1, 1, 357, 1)\n        xs = (xstd - constantMSD.view(1, 1, 357, 1))/constantSSD.view(1, 1, 357, 1)\n        \n        planet_data = (planet_data - constantMEANS)/constantSTDS\n        \n        #x = torch.cat([airs, fgs1], dim=3)  # dim=3 is the 4th dimension (zero-based counting)\n        #print(x.shape, xmean.shape,xstd.shape,xm.shape,xs.shape)  # Should show: torch.Size([1, 2, 5625, 357, 32])\n        \n        x = torch.transpose(x,2,3)\n        x = x.unsqueeze(4)\n        \n        #print(f\"Input: {x.shape}\")\n        for idx, block in enumerate(self.blocks):\n            x = block(x)\n            #print(f\"After Residual Block {idx+1}: {x.shape}\")\n        x = nn.functional.adaptive_avg_pool3d(x, 1)\n        #print(f\"After global avg pool: {x.shape}\")\n        x = x.view(x.size(0), -1)\n\n        #print(\"before cat\",x.shape)\n        ###concatenate params\n        #print(x.shape,xm.shape,xs.shape,planet_data.shape)\n        #x = torch.cat([x,xm.squeeze().unsqueeze(0),xs.squeeze().unsqueeze(0),planet_data.squeeze(1)], dim = 1)\n        x = torch.cat([x,xm.squeeze(),xs.squeeze(),planet_data.squeeze(1)], dim = 1)\n\n        #print(\"after cat\",x.shape)\n\n        \n        x = F.relu(self.fc1(x))\n        #print(f\"After shared FC1: {x.shape}\")\n        indices = [0] + list(range(39, 321))\n        xstd = xstd[:, 0, indices, 0]\n        xmean = xmean[:, 0, indices, 0]\n        pred_x = (self.fc_x(x)*xstd)+ xmean\n        pred_y = (self.fc_y(x)*xstd)+ xmean\n        #print(f\"After prediction head: {pred_y.shape}\")\n        log_sigma = self.fc_sigma(x)\n        #print(f\"After sigma head: {log_sigma.shape}\")\n        return (pred_y-pred_x)/pred_y, log_sigma\n\nstar_info_df = pd.read_csv(\"/kaggle/input/ariel-data-challenge-2025/test_star_info.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:48:18.013255Z","iopub.execute_input":"2025-09-24T23:48:18.013476Z","iopub.status.idle":"2025-09-24T23:48:18.483104Z","shell.execute_reply.started":"2025-09-24T23:48:18.013459Z","shell.execute_reply":"2025-09-24T23:48:18.482512Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\nsample_submission = pd.read_csv(\"/kaggle/input/ariel-data-challenge-2025/sample_submission.csv\")  # or .txt if appropriate\ncolumns = sample_submission.columns.tolist()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:48:20.744919Z","iopub.execute_input":"2025-09-24T23:48:20.745604Z","iopub.status.idle":"2025-09-24T23:48:20.764518Z","shell.execute_reply.started":"2025-09-24T23:48:20.745581Z","shell.execute_reply":"2025-09-24T23:48:20.764006Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def weighted_avg_and_uncertainty(grp):\n    preds = grp[wl_cols].to_numpy(dtype=float)\n    sigmas = grp[sigma_cols].to_numpy(dtype=float)\n    weights = 1.0 / (sigmas ** 2)\n    weighted_preds = np.sum(preds * weights, axis=0) / np.sum(weights, axis=0)\n    sigma_new = 1.0 / np.sqrt(np.sum(weights, axis=0))\n    result = pd.Series(np.concatenate([weighted_preds, sigma_new]))\n    return result\n\ndef mean_of_smallest_sigma_rows(group, wl_cols, sigma_cols, n=10):\n    # Explicitly drop grouping columns to avoid warning\n    group_values = group[wl_cols + sigma_cols].copy()\n    group_values['row_sigma'] = group_values[sigma_cols].mean(axis=1)\n    smallest = group_values.nsmallest(n, 'row_sigma')\n    return smallest[wl_cols + sigma_cols].mean()\n\ndef weighted_avg_and_uncertainty_ii(grp):\n    preds = grp[wl_cols].to_numpy(dtype=float)\n    sigmas = grp[sigma_cols].to_numpy(dtype=float)\n\n    weights = 1.0 / (sigmas ** 2)\n    weighted_preds = np.sum(preds * weights, axis=0) / np.sum(weights, axis=0)\n\n    # Basic average of sigmas, not weighted\n    sigma_new = np.exp(np.mean(np.log(sigmas), axis=0))\n\n    result = pd.Series(np.concatenate([weighted_preds, sigma_new]))\n    return result\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:48:21.212536Z","iopub.execute_input":"2025-09-24T23:48:21.212901Z","iopub.status.idle":"2025-09-24T23:48:21.219614Z","shell.execute_reply.started":"2025-09-24T23:48:21.212874Z","shell.execute_reply":"2025-09-24T23:48:21.219027Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"start_loading = time.perf_counter()\n\nrows = []\nfor index_chunk in  index_chunks:\n    AIRS_CH0_clean = np.ma.MaskedArray(np.zeros((CHUNKS_SIZE, 11250, 32, l)))\n    FGS1_clean = np.ma.MaskedArray(np.zeros((CHUNKS_SIZE, 135000, 32, 32)))\n    \n    chunk_name = '__'.join([f\"{pid}_{obs}\" for pid, obs in index_chunk])\n    \n    if already_saved(chunk_name, path_out):\n            print(f\"Skipping {chunk_name} (already processed)\")\n            continue  # Go to next chunk\n    print(chunk_name)\n    \n    for i in range (CHUNKS_SIZE) : \n        df = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/AIRS-CH0_signal_{index_chunk[i][1]}.parquet'))\n        signal = df.values.astype(np.float64).reshape((df.shape[0], 32, 356))\n        gain = train_adc_info['AIRS-CH0_adc_gain'][0]\n        offset = train_adc_info['AIRS-CH0_adc_offset'][0]\n        signal = ADC_convert(signal, gain, offset)\n        dt_airs = axis_info['AIRS-CH0-integration_time'].dropna().values\n        dt_airs[1::2] += 0.1\n        chopped_signal = signal[:, :, cut_inf:cut_sup]\n        del signal, df\n        \n        # CLEANING THE DATA: AIRS\n        flat = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/AIRS-CH0_calibration_{index_chunk[i][1]}/flat.parquet')).values.astype(np.float64).reshape((32, 356))[:, cut_inf:cut_sup]\n        dark = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/AIRS-CH0_calibration_{index_chunk[i][1]}/dark.parquet')).values.astype(np.float64).reshape((32, 356))[:, cut_inf:cut_sup]\n        dead_airs = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/AIRS-CH0_calibration_{index_chunk[i][1]}/dead.parquet')).values.astype(np.float64).reshape((32, 356))[:, cut_inf:cut_sup]\n        linear_corr = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/AIRS-CH0_calibration_{index_chunk[i][1]}/linear_corr.parquet')).values.astype(np.float64).reshape((6, 32, 356))[:, :, cut_inf:cut_sup]\n        \n        if DO_MASK:\n            chopped_signal = mask_hot_dead(chopped_signal, dead_airs, dark)\n            AIRS_CH0_clean[i] = chopped_signal\n        else:\n            AIRS_CH0_clean[i] = chopped_signal\n            \n        if DO_THE_NL_CORR: \n            linear_corr_signal = apply_linear_corr(linear_corr,AIRS_CH0_clean[i])\n            AIRS_CH0_clean[i,:, :, :] = linear_corr_signal\n        del linear_corr\n        \n        if DO_DARK: \n            cleaned_signal = clean_dark(AIRS_CH0_clean[i], dead_airs, dark, dt_airs)\n            AIRS_CH0_clean[i] = cleaned_signal\n        else: \n            pass\n        del dark\n        \n        df = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/FGS1_signal_{index_chunk[i][1]}.parquet'))\n        fgs_signal = df.values.astype(np.float64).reshape((df.shape[0], 32, 32))\n        \n        FGS1_gain = train_adc_info['FGS1_adc_gain'][0]\n        FGS1_offset = train_adc_info['FGS1_adc_offset'][0]\n        \n        fgs_signal = ADC_convert(fgs_signal, FGS1_gain, FGS1_offset)\n        dt_fgs1 = np.ones(len(fgs_signal))*0.1\n        dt_fgs1[1::2] += 0.1\n        chopped_FGS1 = fgs_signal\n        del fgs_signal, df\n        \n        # CLEANING THE DATA: FGS1\n        flat = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/FGS1_calibration_{index_chunk[i][1]}/flat.parquet')).values.astype(np.float64).reshape((32, 32))\n        dark = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/FGS1_calibration_{index_chunk[i][1]}/dark.parquet')).values.astype(np.float64).reshape((32, 32))\n        dead_fgs1 = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/FGS1_calibration_{index_chunk[i][1]}/dead.parquet')).values.astype(np.float64).reshape((32, 32))\n        linear_corr = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/FGS1_calibration_{index_chunk[i][1]}/linear_corr.parquet')).values.astype(np.float64).reshape((6, 32, 32))\n        \n        if DO_MASK:\n            chopped_FGS1 = mask_hot_dead(chopped_FGS1, dead_fgs1, dark)\n            FGS1_clean[i] = chopped_FGS1\n        else:\n            FGS1_clean[i] = chopped_FGS1\n\n        if DO_THE_NL_CORR: \n            linear_corr_signal = apply_linear_corr(linear_corr,FGS1_clean[i])\n            FGS1_clean[i,:, :, :] = linear_corr_signal\n        del linear_corr\n        \n        if DO_DARK: \n            cleaned_signal = clean_dark(FGS1_clean[i], dead_fgs1, dark,dt_fgs1)\n            FGS1_clean[i] = cleaned_signal\n        else: \n            pass\n        del dark\n        \n    # SAVE DATA AND FREE SPACE\n    AIRS_cds = get_cds(AIRS_CH0_clean)\n    FGS1_cds = get_cds(FGS1_clean)\n\n    del AIRS_CH0_clean, FGS1_clean\n\n    if FILT:\n        AIRS_cds = median_filter_time(AIRS_cds)\n        FGS1_cds = median_filter_time(FGS1_cds)\n    \n    ## (Optional) Time Binning to reduce space\n    if TIME_BINNING:\n        AIRS_cds_binned = bin_obs(AIRS_cds,binning=1)\n        FGS1_cds_binned = bin_obs(FGS1_cds,binning=12*1)\n    else:\n        #AIRS_cds = AIRS_cds.transpose(0,1,3,2) ## this is important to make it consistent for flat fielding, but you can always change it\n        AIRS_cds_binned = AIRS_cds\n        #FGS1_cds = FGS1_cds.transpose(0,1,3,2)\n        FGS1_cds_binned = FGS1_cds\n    AIRS_cds_binned = AIRS_cds_binned.transpose(0,1,3,2)\n    FGS1_cds_binned = FGS1_cds_binned.transpose(0,1,3,2)\n    del AIRS_cds, FGS1_cds\n    \n    for i in range (CHUNKS_SIZE):\n        flat_airs = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/AIRS-CH0_calibration_{index_chunk[i][1]}/flat.parquet')).values.astype(np.float64).reshape((32, 356))[:, cut_inf:cut_sup]\n        flat_fgs = pd.read_parquet(os.path.join(path_folder,f'test/{index_chunk[i][0]}/FGS1_calibration_{index_chunk[i][1]}/flat.parquet')).values.astype(np.float64).reshape((32, 32))\n        if DO_FLAT:\n            corrected_AIRS_cds_binned = correct_flat_field(flat_airs,dead_airs, AIRS_cds_binned[i])\n            AIRS_cds_binned[i] = corrected_AIRS_cds_binned\n            corrected_FGS1_cds_binned = correct_flat_field(flat_fgs,dead_fgs1, FGS1_cds_binned[i])\n            FGS1_cds_binned[i] = corrected_FGS1_cds_binned\n        else:\n            pass\n\n    AIRS_cds_binned = AIRS_cds_binned.transpose(0,1,3,2)\n    FGS1_cds_binned = FGS1_cds_binned.transpose(0,1,3,2)\n    \n    # Example: FGS1_cds (shape: [time, x, y]) -- inpaint along time as channels\n    # Suppose you have a masked array: FGS1_cds (time, x, y), with mask True for bad voxels\n    \n    # Convert to plain array and mask for inpainting\n    data = FGS1_cds_binned[0,:,:,:].data         # shape: (time, x, y)\n    mask = FGS1_cds_binned[0,0,:,:].mask         # shape: (x, y)\n    data = data.transpose(1,2,0)                 # shape: (x, y, time)\n    nan_mask = np.sum(np.isnan(data))\n    if nan_mask:\n        print(\"data contains nan NANANANANANANANANANA\")\n   \n    # Inpaint, treating time as channels (axis=0)\n    result_fgs1 = inpaint_biharmonic(data, mask, channel_axis=2)\n    result_fgs1 = result_fgs1.transpose(2,0,1)\n\n    data_airs = AIRS_cds_binned[0,:,:,:].data      # shape: (time, x, lambda)\n    mask_airs = AIRS_cds_binned[0,0,:,:].mask      # shape: (x, lambda)\n    data_airs = data_airs.transpose(1,2,0)          # shape: (x, lambda, time)\n    # Inpaint, treating wavelength as channels (axis=0)\n    result_airs = inpaint_biharmonic(data_airs, mask_airs, channel_axis=2)\n    result_airs = result_airs.transpose(2,0,1)\n\n    #data_3d_airs = torch.from_numpy(result_airs)      # shape: [frames, x, y]\n    #mask_2d_airs = torch.from_numpy(AIRS_cds_binned.mask[0,0,:,:])      # shape: [x, y]\n    #data_3d_fgs1 = torch.from_numpy(result_fgs1)      # shape: [frames, x, y]\n    #mask_2d_fgs1 = torch.from_numpy(FGS1_cds_binned.mask[0,0,:,:])      # shape: [x, y]\n\n    #sum spatial dimension\n    result_fgs1 = np.sum(result_fgs1, axis=(1, 2))\n    result_airs = np.sum(result_airs, axis=1)\n    #print(result_fgs1.shape,result_airs.shape)\n\n    xmins = []\n    polys = []\n    datas = []\n    \n    #median filter\n    result_fgs1, xcrop = median_filter_and_downsample(result_fgs1, median_filter_window=101, stride=1, plot=False)\n    #print(result_fgs1.shape,result_airs.shape)\n    #find change points\n    #try:\n    onset, offset, mind, xmin = plot_transit_edges(result_fgs1, plot=False)\n    linear = False\n    #print('success')\n    #fit P\n    fitted_curve, coeffs, idx_baseline = fit_and_plot_baseline(result_fgs1,onset,offset,delta=10,degree=2,planet_id=None,plot=False)\n    datas.append(result_fgs1)\n    xmins.append(xmin)\n    polys.append(fitted_curve)\n    #except:\n        #if not do linear\n        #linear = True\n\n    for wl in range(result_airs.shape[1]):\n        signal = result_airs[:, wl]\n        signal, xcrop = median_filter_and_downsample(signal, median_filter_window=101, stride=1, plot=False)\n        #fitted_curve, coeffs, idx_baseline = fit_and_plot_baseline(\n        #    signal,\n        #    onset,\n        #    offset,\n        #    delta=10,\n        #    degree=2,\n        #    planet_id=None,\n        #    plot=False\n        #)\n        #smoothed_signal = savgol_filter(signal, window_length=15, polyorder=2)\n        datas.append(signal)\n        #xmins.append(smoothed_signal[mind])\n        #polys.append(fitted_curve)\n    \n    #save\n\n    datas_tensor = torch.from_numpy(np.stack(datas))   # Shape: (num_arrays, array_length)\n    polys_tensor = torch.from_numpy(np.stack(polys))   # Shape: (num_arrays, array_length)\n    \n    # Convert list of scalars to 1D tensor\n    xmins_tensor = torch.tensor(xmins)                  # Shape: (num_scalars,)\n    \n    #torch.save({'data': datas_tensor, 'poly': polys_tensor, 'xmin': xmins_tensor, 'mind':torch.tensor(mind)}, os.path.join(path_out, f'clean_train_{chunk_name}.pt'))\n    \n    \n    model = Custom2DCNN(\n        in_channels=1,\n        conv_kernel_sizes=[3,3,3,3,3,3],\n        conv_filters=[16,32,64,128,256,512],\n        pool_kernel_sizes=[[4,2,1],[4,2,1],[4,2,1],[4,2,1],[2,2,1],[2,2,1]],\n        pool_strides=[[2,1,1],[2,1,1],[2,1,1],[2,1,1],[2,1,1],[2,1,1]],\n        dilation_rates=[[1,1,1], [1,3,1], [1,9,1], [1,27,1], [1,27*3,1], [1,27*9,1]],\n        circular_pad_wavelength_layers=[True, True, True, True, True, True]\n    ).to(device)\n    checkpoint = torch.load(\"/kaggle/input/runc1_0-best/pytorch/default/1/best_model (1).pth\", weights_only=False).to(device)\n    model = checkpoint\n\n    offsets = torch.arange(5)  # tensor([0,1,2,3,4])\n    stride = 10\n    slices = [datas_tensor[:, offset::stride] for offset in offsets]  # list of 5 tensors, shape [wavelength, time_downsampled]\n    # Get the flipped (time domain reversed) versions of these slices\n    #slices_flipped = [torch.flip(s, dims=[1]) for s in slices]  # assume time dimension is dim=1\n    \n    #all_slices = slices + slices_flipped\n    #print(all_slices)\n    # Stack into batch dimension (dim=0)\n    data = torch.stack(slices, dim=0) \n    \n    planet_data = torch.tensor(get_planet_values(star_info_df, [index_chunk[i][0]]))\n    #data = data.unsqueeze(0) \n    data = data.unsqueeze(1) \n    planet_data.unsqueeze(0)\n    planet_data = planet_data.repeat(5, 1)\n    model.eval()\n    with torch.no_grad():\n        y_pred, log_sigma = model(data.float().to(device), planet_data.float().to(device))\n    #print(y_pred.shape,log_sigma.shape)\n\n    y_pred[y_pred < 0] = 0\n\n    batch_size = y_pred.shape[0]  # e.g., 5\n    planet_id = index_chunk[i][0]  # Same planet_id for entire batch\n    \n    for batch_i in range(batch_size):\n        row_dict = {}\n        row_dict['planet_id'] = planet_id  # Same for all batch items\n        sigma_exp = torch.exp(log_sigma)\n        for j in range(1, 284):\n            row_dict[f'wl_{j}'] = y_pred[batch_i, j-1].item()\n            row_dict[f'sigma_{j}'] = sigma_exp[batch_i, j-1].item()\n        #for k in range(7):\n        rows.append(row_dict)\n\n    checkpoint = torch.load(\"/kaggle/input/c1f3/pytorch/default/1/best_model.pth\", weights_only=False).to(device)\n    model = checkpoint\n    model.eval()\n    with torch.no_grad():\n        y_pred, log_sigma = model(data.float().to(device), planet_data.float().to(device))\n    #print(y_pred.shape,log_sigma.shape)\n\n    y_pred[y_pred < 0] = 0\n\n    batch_size = y_pred.shape[0]  # e.g., 5\n    planet_id = index_chunk[i][0]  # Same planet_id for entire batch\n    \n    for batch_i in range(batch_size):\n        row_dict = {}\n        row_dict['planet_id'] = planet_id  # Same for all batch items\n        sigma_exp = torch.exp(log_sigma)\n        for j in range(1, 284):\n            row_dict[f'wl_{j}'] = y_pred[batch_i, j-1].item()\n            row_dict[f'sigma_{j}'] = sigma_exp[batch_i, j-1].item()\n        rows.append(row_dict)\n    \n    \n    #print(chunk_name, count)\n    del AIRS_cds_binned\n    del FGS1_cds_binned\n    count +=1\n\nsubmission_df = pd.DataFrame(rows, columns=columns)\n\n### BEST SIGMA\n# Find index of row with minimum average sigma for each planet_id\n#submission_df['mean_sigma'] = submission_df[[f'sigma_{i}' for i in range(1, 284)]].mean(axis=1)\n#df_best = submission_df.loc[submission_df.groupby('planet_id')['mean_sigma'].idxmin()].drop(columns='mean_sigma')\n\n### BASIC AVERAGE\n# List of columns for predictions and sigmas\nwl_cols = [f'wl_{i}' for i in range(1, 284)]\nsigma_cols = [f'sigma_{i}' for i in range(1, 284)]\n\n# Take simple mean for each group (planet_id)\ndf_best = (\n    submission_df.groupby('planet_id')[wl_cols + sigma_cols]\n    .mean()\n    .reset_index()\n)\n\n### BEST N AVERAGE\n#df_best = (\n#    submission_df.groupby('planet_id')\n#    .apply(mean_of_smallest_sigma_rows, wl_cols=wl_cols, sigma_cols=sigma_cols, n=10)\n#    .reset_index()\n#)\n\n### FANCY AVERAGE (II)\n# Get column names for output\n#output_cols = wl_cols + sigma_cols\n\n#df_best = (\n#    submission_df.groupby('planet_id')[wl_cols + sigma_cols].apply(weighted_avg_and_uncertainty_ii)\n#    .reset_index()\n#)\n#df_best.columns = ['planet_id'] + output_cols\n\n\ndf_best.to_csv(\"submission.csv\", index=False)\n\nloading_time = time.perf_counter() - start_loading\nprint(f\"Start processing data time: {loading_time:.4f} sec\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:54:05.113568Z","iopub.execute_input":"2025-09-24T23:54:05.114317Z","iopub.status.idle":"2025-09-24T23:54:35.204395Z","shell.execute_reply.started":"2025-09-24T23:54:05.114294Z","shell.execute_reply":"2025-09-24T23:54:35.203573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#torch.save({'data': torch.from_numpy(result_airs)}, os.path.join(path_out, f'AIRS_clean_train_{chunk_name}.pt'))\ndf_best","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:49:01.089487Z","iopub.execute_input":"2025-09-24T23:49:01.089741Z","iopub.status.idle":"2025-09-24T23:49:01.114008Z","shell.execute_reply.started":"2025-09-24T23:49:01.089722Z","shell.execute_reply":"2025-09-24T23:49:01.113376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#submission_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T23:54:35.205774Z","iopub.execute_input":"2025-09-24T23:54:35.206003Z","iopub.status.idle":"2025-09-24T23:54:35.228827Z","shell.execute_reply.started":"2025-09-24T23:54:35.205984Z","shell.execute_reply":"2025-09-24T23:54:35.228024Z"}},"outputs":[],"execution_count":null}]}