{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.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":70367,"databundleVersionId":9188054,"sourceType":"competition"}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Overview\nBased on official calibration notebooks (https://www.kaggle.com/code/gordonyip/update-calibrating-and-binning-astronomical-data, https://www.kaggle.com/code/gordonyip/calibrating-a-single-observation/notebook), I reimplemented calibration with JAX.\n\n\nAccoding to a discussion https://www.kaggle.com/competitions/ariel-data-challenge-2024/discussion/528247, the calibration preprocess takes too much time and it is necessary to speed up to fit 9h time limit.\n\n\nI also added feature extraction of transit depth (aka. intensity reduction ratio),\nand it is still faster than the original time described above.\n\n\n# Differences\n- Don't iterate 5 times to identify 'hot' pixel. (I'm not sure whether it is important)\n- Use different time binning of 25 (instead of 30), since the number of records cannot be divided by 30.\n\n\n# Future Works\n- Investigate feature extraction. I'm not sure whether this simple algorithm is fine.\n- Think how to estimate uncertainty. (Should we consider the number of dead/hot pixels?)","metadata":{}},{"cell_type":"code","source":"import functools\nimport itertools\n\nimport numpy as np\nimport pandas as pd\nimport jax\nimport jax.numpy as jnp\nfrom tqdm.notebook import tqdm, trange\nimport matplotlib.pyplot as plt","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-08-24T01:15:06.429647Z","iopub.execute_input":"2024-08-24T01:15:06.430073Z","iopub.status.idle":"2024-08-24T01:15:09.219855Z","shell.execute_reply.started":"2024-08-24T01:15:06.430031Z","shell.execute_reply":"2024-08-24T01:15:09.218579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FOLDER = '/kaggle/input/ariel-data-challenge-2024/'","metadata":{"execution":{"iopub.status.busy":"2024-08-24T01:15:09.222658Z","iopub.execute_input":"2024-08-24T01:15:09.223861Z","iopub.status.idle":"2024-08-24T01:15:09.231516Z","shell.execute_reply.started":"2024-08-24T01:15:09.223813Z","shell.execute_reply":"2024-08-24T01:15:09.228725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_adc_info = pd.read_csv(f\"{FOLDER}/train_adc_info.csv\", index_col = \"planet_id\")\ndisplay(train_adc_info)","metadata":{"execution":{"iopub.status.busy":"2024-08-24T01:15:09.237677Z","iopub.execute_input":"2024-08-24T01:15:09.238036Z","iopub.status.idle":"2024-08-24T01:15:09.288745Z","shell.execute_reply.started":"2024-08-24T01:15:09.238Z","shell.execute_reply":"2024-08-24T01:15:09.287588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_adc_info = pd.read_csv(f\"{FOLDER}/test_adc_info.csv\", index_col = \"planet_id\")\ndisplay(test_adc_info)","metadata":{"execution":{"iopub.status.busy":"2024-08-24T01:15:09.291747Z","iopub.execute_input":"2024-08-24T01:15:09.292522Z","iopub.status.idle":"2024-08-24T01:15:09.311422Z","shell.execute_reply.started":"2024-08-24T01:15:09.292478Z","shell.execute_reply":"2024-08-24T01:15:09.309797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"axis_info = pd.read_parquet(f\"{FOLDER}/axis_info.parquet\")\ndisplay(axis_info)","metadata":{"execution":{"iopub.status.busy":"2024-08-24T01:15:09.313626Z","iopub.execute_input":"2024-08-24T01:15:09.314115Z","iopub.status.idle":"2024-08-24T01:15:09.535279Z","shell.execute_reply.started":"2024-08-24T01:15:09.314072Z","shell.execute_reply":"2024-08-24T01:15:09.533848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = pd.read_csv(f\"{FOLDER}/train_labels.csv\", index_col=\"planet_id\")\ndisplay(train_labels)","metadata":{"execution":{"iopub.status.busy":"2024-08-24T01:15:09.537639Z","iopub.execute_input":"2024-08-24T01:15:09.53828Z","iopub.status.idle":"2024-08-24T01:15:09.660956Z","shell.execute_reply.started":"2024-08-24T01:15:09.538227Z","shell.execute_reply":"2024-08-24T01:15:09.659787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wavelength = pd.read_csv(f\"{FOLDER}/wavelengths.csv\")\ndisplay(wavelength)","metadata":{"execution":{"iopub.status.busy":"2024-08-24T01:15:09.662313Z","iopub.execute_input":"2024-08-24T01:15:09.662638Z","iopub.status.idle":"2024-08-24T01:15:09.701846Z","shell.execute_reply.started":"2024-08-24T01:15:09.66261Z","shell.execute_reply":"2024-08-24T01:15:09.700627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\ndef check(name, arr):\n    if False:\n        print(f\"\\n{name}: shape: {arr.shape}\")\n        display(arr)\n\ndef read_detector(mode: str, image_id: int, channel: str, name: str, inf: int, sup: int) -> jax.Array:\n    arr = jnp.array(pd.read_parquet(f\"{FOLDER}/{mode}/{image_id}/{channel}_calibration/{name}.parquet\").values[:, inf:sup])\n    check(name, arr)\n    return arr\n\n@jax.jit\ndef ADC(signal: jax.Array, params: jax.Array) -> jax.Array:\n    signal = signal.at[:].set(jnp.divide(signal, params.at[0].get())) # signal /= gain\n    signal = signal.at[:].set(jnp.add(   signal, params.at[1].get())) # signal += offset\n    return signal\n\n@jax.jit\ndef hot_dead_mask(dark: jax.Array, dead: jax.Array) -> jax.Array:\n    # Should I iterate 5times as `astropy.stat.sigma_clip()`?\n    m = dark.mean()\n    five_sigma = jnp.multiply(dark.std(), 5)\n    hot = jnp.logical_or(\n        jnp.greater(m, jnp.add(m, five_sigma)),\n        jnp.less(m, jnp.subtract(m, five_sigma)),\n    )\n\n    return jnp.logical_or(hot, dead)\n\n@jax.jit\ndef adjust_linearity(linear_corr: jax.Array, signal: jax.Array) -> jax.Array:\n    # [0] + [1]*x + [2]*x^2 + [3]*x^3 + [4]*x^4 + [5]*x^5\n    assert linear_corr.shape[1:] == signal.shape[1:], f\"BUG: {linear_corr.shape} vs {signal.shape}\"\n    \n    signal, _ = jax.lax.scan(\n        lambda carry, p: (jnp.add(jnp.multiply(carry, signal), p), None),\n        jnp.zeros_like(signal),\n        linear_corr,\n        reverse=True,\n    )\n    return signal\n\n@jax.jit\ndef subtract_dark(signal: jax.Array, mask: jax.Array, dark: jax.Array, dt: jax.Array) -> jax.Array:\n    return jnp.where(\n        mask,\n        signal,\n        jnp.subtract(signal, jnp.multiply(dark, dt.reshape((-1, 1, 1))))\n    )\n\n@jax.jit\ndef CDS(signal: jax.Array) -> jax.Array:\n    return signal.at[1::2, :, :].get() - signal.at[::2, :, :].get()\n\n\n@jax.jit\ndef correct_flat(signal: jax.Array, mask: jax.Array, flat: jax.Array) -> jax.Array:\n    return jnp.where(mask, signal, jnp.divide(signal, flat))    \n\n@functools.partial(jax.jit, static_argnums=(1,))\ndef smooth_average(signal: jax.Array, binw: int) -> jax.Array:\n    return jax.vmap(\n        lambda idx: jnp.nanmean(jax.lax.dynamic_slice(signal, (idx * binw, 0, 0), (binw, *signal.shape[1:])), axis=0)\n    )(jnp.arange(signal.shape[0] // binw))\n\n\n@functools.partial(jax.jit, static_argnums=(7, 8, 9))\ndef calibrate_impl(signal: jax.Array, dark: jax.Array, dead: jax.Array, flat: jax.Array,\n                   adc_params: jax.Array, linear_corr: jax.Array, dt: jax.Array,\n                   binw: int, inf: int, sup: int) -> jax.Array:\n    mask = hot_dead_mask(dark, dead)\n    \n    signal = signal.at[:].set(ADC(signal, adc_params))\n    signal = signal.at[:, :, inf:sup].get()\n    signal = signal.at[:].set(adjust_linearity(linear_corr, signal))\n    signal = signal.at[:].set(subtract_dark(signal, mask, dark, dt))\n    signal = CDS(signal)\n    signal = signal.at[:].set(correct_flat(signal, mask, flat))\n    signal = smooth_average(signal, binw)\n    signal = signal.at[:].set(jnp.where(mask, jnp.nan, signal))\n\n    return signal.transpose(0, 2, 1) # [time, wavelength/space?, space]\n\n\ndef calibrate(mode: str, image_id: int, channel: str, inf: int, sup: int, binw: int) -> jax.Array:\n    #print(f\"{image_id=}, {channel=}\")\n\n    flat = read_detector(mode, image_id, channel, \"flat\", inf, sup)\n    dark = read_detector(mode, image_id, channel, \"dark\", inf, sup)\n    dead = read_detector(mode, image_id, channel, \"dead\", inf, sup)\n    flat = read_detector(mode, image_id, channel, \"flat\", inf, sup)\n    \n    linear_corr = jnp.array(pd.read_parquet(f\"{FOLDER}/{mode}/{image_id}/{channel}_calibration/linear_corr.parquet\").values.reshape((6, 32, -1))[:, :, inf:sup])\n    check(\"linear_corr\", linear_corr)\n\n    s = pd.read_parquet(f\"{FOLDER}/{mode}/{image_id}/{channel}_signal.parquet\")\n    signal = jnp.array(s.values.reshape((s.shape[0] ,32, -1)), dtype=jnp.float32)\n    check(\"raw signal\", signal)\n\n    adc_params = jnp.array((train_adc_info if mode == \"train\" else test_adc_info).loc[image_id, [f\"{channel}_adc_gain\", f\"{channel}_adc_offset\"]])\n    check(\"gain & offset\", adc_params)\n\n    dt = (jnp.array(axis_info[f'{channel}-integration_time'].dropna().values)\n          if f\"{channel}-integragion_time\" in axis_info.columns\n          else jnp.full((signal.shape[0],), 0.1))\n    check(\"dt\", dt)\n\n    return calibrate_impl(signal, dark, dead, flat, adc_params, linear_corr, dt, binw, inf, sup)\n\n@functools.partial(jax.jit, static_argnums=(1,))\ndef diff(signal: jax.Array, w: int) -> jax.Array:\n    assert signal.ndim == 2, f\"BUG: {signal.ndim}\"\n    return jax.vmap(\n        lambda idx: jnp.subtract(\n            jnp.nanmean(jax.lax.dynamic_slice(signal, (idx+w, 0), (w, signal.shape[1])), axis=0),\n            jnp.nanmean(jax.lax.dynamic_slice(signal, (idx  , 0), (w, signal.shape[1])), axis=0),\n        )\n    )(jnp.arange(signal.shape[0] - 2*w))\n\n@functools.partial(jax.jit, static_argnums=(1,))\ndef transit_depth(signal: jax.Array, w: int) -> jax.Array:\n    d = diff(signal, w)\n\n    dI = jnp.divide(jnp.subtract(jnp.nanmax(d, axis=0), jnp.nanmin(d, axis=0)), 2)\n    idx = jnp.floor_divide(jnp.add(jnp.nanargmax(d, axis=0), jnp.nanargmin(d, axis=0)), 2)\n\n    Idec = jax.vmap(\n        lambda i, s: jnp.nanmean(jax.lax.dynamic_slice(s, (i-w,), (2*w,)), axis=0),\n        in_axes=(0, 1),\n    )(idx, signal)\n\n    return jnp.divide(dI, jnp.add(Idec, dI))\n\n\ndef spectra(mode: str, image_id: int, plot: bool) -> jax.Array:\n    # Calibration\n    AIRS = calibrate(mode, image_id, \"AIRS-CH0\", 39, 321, 25)\n    FGS1 = calibrate(mode, image_id, \"FGS1\"    ,  0,  32, 25*12)\n    assert AIRS.shape[0] == FGS1.shape[0] == 225, f\"BUG: {AIRS.shape[0]}, {FGS1.shape[0]}\"\n\n\n    # Feature Extraction of Transit Depth\n    window = 10\n    AIRS = jnp.nanmean(AIRS, axis=2) # [time, wavelength]\n    FGS1 = jnp.reshape(jnp.nanmean(FGS1, axis=(1, 2)), (-1, 1)) # [time, 1]\n\n    s = (\n        jnp.zeros((283,))\n        .at[0:1].set(transit_depth(FGS1, window))\n        .at[1: ].set(transit_depth(AIRS, window))\n    )\n\n    if plot:\n        plt.plot(jnp.nanmean(AIRS, axis=1), label=\"signal\")\n        plt.twinx().plot(jnp.add(jnp.arange(AIRS.shape[0]-2*window), window),\n                         jnp.reshape(diff(jnp.nanmean(AIRS, axis=1, keepdims=True), window), (-1,)),\n                         color=\"tab:orange\", label=\"diff\")\n        plt.title(\"AIRS-CH0\")\n        plt.legend()\n        plt.show()\n\n        plt.plot(FGS1, label=\"signal\")\n        plt.twinx().plot(jnp.add(jnp.arange(FGS1.shape[0]-2*window), window),\n                         jnp.reshape(diff(FGS1, window), (-1,)),\n                         color=\"tab:orange\", label=\"diff\")\n        plt.title(\"FGS1\")\n        plt.legend()\n        plt.show()\n\n        wave = wavelength.iloc[0].values\n        plt.plot(wave, s, label=\"calcurated\", marker=\".\", linestyle=\"\")\n        plt.plot(wave, train_labels.loc[image_id, :].values, label=\"truth\", marker=\".\", linestyle=\"\")\n        plt.legend()\n        plt.title(\"transit depth\")\n        plt.show()\n\n    return s\n\ndef run():\n    #n = 3 # for Debugging\n    n = train_adc_info.shape[0]\n    \n    calcurated = pd.DataFrame(\n        itertools.islice((np.asarray(spectra(\"train\", train_adc_info.index[i], i==0)) for i in trange(train_adc_info.shape[0])), n),\n        columns=[f\"wl_{i+1}\" for i in range(283)],\n        index=train_adc_info.index[:n],\n    )\n    display(calcurated)\n    \n    calcurated.to_csv(\"train_calcurated.csv\", index=True, header=True)\n\n    error = calcurated - train_labels.loc[calcurated.index]\n    plt.errorbar(wavelength.iloc[0].values, error.mean(axis=0), yerr=error.std(axis=0), marker=\".\", linestyle=\"\")\n    plt.title(\"calcurated - ground truth: mean +/- std\")\n    plt.show()\n\n\nrun()","metadata":{"execution":{"iopub.status.busy":"2024-08-24T01:15:09.70344Z","iopub.execute_input":"2024-08-24T01:15:09.703857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}