{"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":"none","dataSources":[{"sourceId":101849,"databundleVersionId":13093295,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"***Introduction***\n\nIn this notebook, we explored the idea of fitting the out-of-transit portions of each light curve to model and normalize the background signal. Our goal was to improve the reliability of transit depth measurements by accounting for background trends and systematics. However, we ultimately decided not to pursue this approach further for this challenge. Many of the light curves were noisy, exhibited unpredictable background changes, or were only partially observed, making it difficult to robustly disentangle the background from the transit signal.","metadata":{}},{"cell_type":"markdown","source":"We first identify the start and end points of planetary transits. We use a Savitzky-Golay filter to reduce noise, then locate the global minimum (likely the transit midpoint). We split the signal at this minimum and apply binary segmentation (via the ruptures package) to each half to detect where the flux drops (transit onset) and rises (transit offset). These edges are further refined by searching for regions where the signal remains above a certain percentile threshold, ensuring the detected boundaries correspond to real transitions in the light curve.\n\nAfter identifying the transit region, we fit a polynomial baseline to the out-of-transit data to model and normalize the background signal. This method is robust for clean, full transits, but we found it less reliable for partial or noisy signals, so we ultimately decided not to use baseline normalization for all cases. We briefly attempted to fit a BATMAN transit model to overcome the limitations of partial transits, but did not pursue this further due to time constraints and the need for careful validation.\n\nFinally, we compare our estimated transit depths to the provided ground truth values for each curve. We observed that the ground truth transit depth does not consistently correspond to the lowest point in the light curve, nor to a fixed fraction of the fitted transit model. In some cases, the ground truth value is even lower than the minimum observed flux. This variability suggests that the true transit depth depends not only on the observed light curve, but also on additional parameters such as the star and planet's properties and the transit conditions.","metadata":{}},{"cell_type":"markdown","source":"***Resources***\n\nThis notebook uses the ruptures library:\n\nC. Truong, L. Oudre, N. Vayatis. Selective review of offline change point detection methods. Signal Processing, 167:107299, 2020.\n\n\nThis notebook also attempts to fit the transit model using the batman package:\n\nKreidberg, L. (2015). \"batman: BAsic Transit Model cAlculatioN in Python.\" Publications of the Astronomical Society of the Pacific, 127(957), 1161–1165.\n\nThis notebook borrows many preprocessing functions and steps from: https://www.kaggle.com/code/gordonyip/update-calibrating-and-binning-astronomical-data","metadata":{}},{"cell_type":"code","source":"!pip install ruptures\n!pip install ldtk","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-30T19:58:05.479768Z","iopub.execute_input":"2025-09-30T19:58:05.480127Z","iopub.status.idle":"2025-09-30T19:58:14.061455Z","shell.execute_reply.started":"2025-09-30T19:58:05.4801Z","shell.execute_reply":"2025-09-30T19:58:14.05946Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install batman-package\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-30T19:58:14.063888Z","iopub.execute_input":"2025-09-30T19:58:14.064271Z","iopub.status.idle":"2025-09-30T19:58:18.63482Z","shell.execute_reply.started":"2025-09-30T19:58:14.064237Z","shell.execute_reply":"2025-09-30T19:58:18.633565Z"}},"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\nimport ruptures as rpt\nfrom scipy.signal import savgol_filter\nimport seaborn as sns\nimport batman\nfrom ldtk import LDPSetCreator, BoxcarFilter","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-09-30T19:58:18.637086Z","iopub.execute_input":"2025-09-30T19:58:18.637456Z","iopub.status.idle":"2025-09-30T19:58:20.800875Z","shell.execute_reply.started":"2025-09-30T19:58:18.637415Z","shell.execute_reply":"2025-09-30T19:58:20.799802Z"}},"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(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\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(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 get_cds(signal):\n    cds = signal[:,1::2,:,:] - signal[:,::2,:,:]\n    return cds\n\ndef correct_flat_field(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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-30T19:58:20.8032Z","iopub.execute_input":"2025-09-30T19:58:20.803532Z","iopub.status.idle":"2025-09-30T19:58:20.81414Z","shell.execute_reply.started":"2025-09-30T19:58:20.803508Z","shell.execute_reply":"2025-09-30T19:58:20.812964Z"}},"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 = rpt.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 = rpt.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-30T19:58:20.815644Z","iopub.execute_input":"2025-09-30T19:58:20.816436Z","iopub.status.idle":"2025-09-30T19:58:20.856113Z","shell.execute_reply.started":"2025-09-30T19:58:20.8164Z","shell.execute_reply":"2025-09-30T19:58:20.854868Z"}},"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, 'train','*','*'))\n\nCHUNKS_SIZE = 1\nindex_chunks = get_multiobs_index(files, CHUNKS_SIZE)\n\nmidpoint = len(index_chunks) // 2\nfirst_half_chunks = index_chunks[38:58]\nsecond_half_chunks = index_chunks[midpoint:]\n\nindex_chunks = first_half_chunks\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\ntrain_df = pd.read_csv('/kaggle/input/ariel-data-challenge-2025/train.csv')\nstar_train_df = pd.read_csv('/kaggle/input/ariel-data-challenge-2025/train_star_info.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-30T20:59:18.012576Z","iopub.execute_input":"2025-09-30T20:59:18.01322Z","iopub.status.idle":"2025-09-30T20:59:19.820499Z","shell.execute_reply.started":"2025-09-30T20:59:18.01318Z","shell.execute_reply":"2025-09-30T20:59:19.819646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ratios = []\nstati = []\nstatTs = []\nstatsma = []\nplanets = []\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'train/{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'train/{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'train/{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'train/{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'train/{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'train/{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'train/{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'train/{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'train/{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'train/{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'train/{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'train/{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\")\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=True)\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=True)\n    datas.append(result_fgs1)\n    xmins.append(xmin)\n    polys.append(fitted_curve)\n\n    \n    VAL = train_df.loc[train_df['planet_id'] == index_chunk[i][0], 'wl_1'].values\n    if VAL.size == 0:\n        print(f\"No wl_1 value found for planet_id {index_chunk[i][0]}\")\n        VAL = 1.0  # Or set some default fallback\n    else:\n        VAL = VAL[0]\n\n    # Calculate the modified curve\n    modified_curve = fitted_curve - (fitted_curve * VAL)\n\n    sval = np.max((fitted_curve-result_fgs1)/fitted_curve)\n    ratio = VAL/sval\n    ratios.append(ratio)\n\n    stati.append(star_train_df.loc[star_train_df['planet_id']==index_chunk[i][0],'i'].values[0])\n    statTs.append(star_train_df.loc[star_train_df['planet_id']==index_chunk[i][0],'Ts'].values[0])\n    statsma.append(star_train_df.loc[star_train_df['planet_id']==index_chunk[i][0],'sma'].values[0])\n    planets.append(index_chunk[i][0])\n\n    # Create DataFrame\n    df_stats = pd.DataFrame({\n        'planet_id': planets,\n        'ratio': ratios,\n        'i': stati,\n        'Ts': statTs,\n        'sma': statsma\n    })\n       \n    # result_fgs1: your observed light curve\n    time = (np.arange(len(result_fgs1))-mind)*4*1.2/86400\n    \n    # Use pre-transit baseline to normalize, if enough points; otherwise, use median\n    pre_transit = result_fgs1[:onset] if onset > 0 else result_fgs1\n    flux = result_fgs1 / np.median(pre_transit)\n    \n    # Assign provided parameters for the current planet:\n    period = star_train_df.loc[star_train_df['planet_id']==index_chunk[i][0],'P'].values[0]\n    a_rs = star_train_df.loc[star_train_df['planet_id']==index_chunk[i][0],'sma'].values[0]\n    inc = star_train_df.loc[star_train_df['planet_id']==index_chunk[i][0],'i'].values[0]\n    ecc = star_train_df.loc[star_train_df['planet_id']==index_chunk[i][0],'e'].values[0]\n    t0 = 0#(onset + offset) / 2  # index for mid-transit; if you have real time, use that!\n    rp_rs = np.sqrt(sval)       # if you know R_p, rp_rs = R_p / Rs; else, sqrt(depth)\n    \n    # Quick setup of limb darkening coefficients (here, quadratic, dummy values)\n    #u1, u2 = 0.2, 0.3  # You can use more physical values if you have tables by Ts\n\n    \n    # Retrieve Ts from your dataframe (already formatted per ID in your loop)\n    Ts = float(star_train_df.loc[star_train_df['planet_id'] == index_chunk[i][0], 'Ts'].values[0])\n    Ts_err = 100                # Set uncertainty, or use your own estimates if available\n    logg = 4.5\n    logg_err = 0.1\n    z = 0.0                     # Metallicty\n    z_err = 0.05\n    \n    # Define your filter, e.g. 500-600 nm (adjust to your passband as needed)\n    filters = [BoxcarFilter('myband', 500, 600)]\n    \n    # Generate LD coefficients\n    sc = LDPSetCreator(teff=(Ts, Ts_err), logg=(logg, logg_err), z=(z, z_err), filters=filters)\n    ps = sc.create_profiles()\n    u, uerr = ps.coeffs_qd(do_mc=True)  # This gets quadratic LD coeffs\n    \n    u1, u2 = u[0]\n    print(f\"Limb darkening for Ts={Ts}: u1={u1:.3f}, u2={u2:.3f}\")\n        \n    params = batman.TransitParams()\n    params.t0 = t0\n    params.per = period\n    params.rp = rp_rs\n    params.a = a_rs\n    params.inc = inc\n    params.ecc = ecc\n    params.w = 90.0  # argument of periastron; 90 deg is common for transits\n    params.u = [u1, u2]\n    params.limb_dark = \"quadratic\"\n    \n    m = batman.TransitModel(params, time)\n    model_flux = m.light_curve(params)\n    \n    # Plot\n    plt.figure(figsize=(10,6))\n    plt.plot(time, flux, label='result_fgs1 (normalized)', color='blue')\n    plt.plot(time, fitted_curve / np.median(pre_transit), label='fitted_curve (normalized)', color='orange')\n    plt.plot(time, model_flux, label='BATMAN transit fit', color='green')\n    plt.legend()\n    plt.title(f'Transit Fit for Planet {index_chunk[i][0]}')\n    plt.xlabel('Time Index')\n    plt.ylabel('Normalized Flux')\n    plt.show()\n\n    \n    # Plotting\n    plt.figure(figsize=(10,6))\n    plt.plot(result_fgs1, label='result_fgs1')\n    plt.plot(fitted_curve, label='fitted_curve')\n    plt.plot(modified_curve, label='fitted_curve * VAL')\n    plt.legend()\n    plt.title(f'Transit Curve for Planet {index_chunk[i][0]}')\n    plt.xlabel('Index or Time')\n    plt.ylabel('Signal Strength')\n    plt.show()\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    #torch.save({'data': data_3d_fgs1, 'mask': mask_2d_fgs1}, os.path.join(path_out, f'FGS1_clean_train_{chunk_name}.pt'))\n    \n    print(chunk_name, count)\n    del AIRS_cds_binned\n    del FGS1_cds_binned\n    count +=1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-30T20:59:19.822055Z","iopub.execute_input":"2025-09-30T20:59:19.822348Z","iopub.status.idle":"2025-09-30T21:12:01.537339Z","shell.execute_reply.started":"2025-09-30T20:59:19.822318Z","shell.execute_reply":"2025-09-30T21:12:01.535761Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"For each planet, the notebook generates a series of four figures to illustrate the transit detection and analysis process:\n\n*Transit Detection Edges:*\n\nThis figure shows the beginning (onset) and end (offset) of the detected transit in the light curve.\n\n*Fitted Baseline and Transit Edges:*\n\nHere, the fitted baseline (the expected signal outside of transit) is plotted along with the baseline points used for fitting and the transit onset/offset.\n\n*BATMAN Transit Model Fit:*\n\nThis figure overlays an attempted BATMAN transit model fit on the observed light curve. \n\n*Ground Truth Comparison (fitted_curve, original curve, fitted_curve * VAL):*\n\nThe final figure compares the fitted baseline curve, the original observed curve, and the curve modified by the ground truth transit value (VAL). The line for fitted_curve * VAL is particularly interesting: for some planets, it aligns with the minimum of the transit curve or a certain fraction above the minimum, but for others, it is even lower than the observed minimum. This variability suggests that the ground truth transit depth is not always at the lowest point of the curve and can be influenced by other parameters besides the light curve.","metadata":{}},{"cell_type":"code","source":"# Plotting\nparams = ['i', 'Ts', 'sma']\nfor param in params:\n    plt.figure(figsize=(8, 5))\n    sns.scatterplot(data=df_stats, x=param, y='ratio')\n    plt.title(f'Ratio vs {param}')\n    plt.xlabel(param)\n    plt.ylabel('Ratio (VAL / sval)')\n    plt.grid(True)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-30T20:38:06.573693Z","iopub.execute_input":"2025-09-30T20:38:06.57653Z","iopub.status.idle":"2025-09-30T20:38:07.315312Z","shell.execute_reply.started":"2025-09-30T20:38:06.576457Z","shell.execute_reply":"2025-09-30T20:38:07.314026Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"After processing and fitting the transit curves for each planet, we generate scatter plots showing how the ratio $\\text{VAL} / s_{\\text{val}}$ (where $\\text{VAL}$ is the ground truth transit depth and $s_{\\text{val}}$ is the maximum observed transit depth from our fit) varies with three key system parameters:\n\n- *Inclination (i):* The angle of the planet's orbit relative to our line of sight. Higher inclinations (closer to 90°) mean the planet transits across the center of the star, typically producing deeper and more symmetric transits.\n- *Stellar Temperature (Ts):* The effective temperature of the host star. This can influence limb darkening and the overall shape of the transit.\n- *Semi-major Axis (sma):* The distance between the planet and the star.\n\nThese figures help assess whether any of these physical parameters systematically influence the ratio between the ground truth transit depth and the observed transit depth. ","metadata":{}}]}