{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":101849,"databundleVersionId":13093295},{"sourceType":"datasetVersion","sourceId":13302394,"datasetId":8076891,"databundleVersionId":14004807}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This notebook demonstrates model fitting and boosting inference. Attached dataset already contains checkpoints for boosting. The training of boostings is demonstrated in `ARIEL-2025-boosting` notebook: https://www.kaggle.com/code/alehandreus/ariel-2025-boosting.\n\nThis version does inference on first 10 train planets and calculates the score. For hidden-test-ready variant, see `test` version.","metadata":{}},{"cell_type":"code","source":"%%writefile median_pool.py\n\n\n# https://gist.github.com/rwightman/f2d3849281624be7c0f11c85c87c1598\n\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn.modules.utils import _pair, _quadruple\n\n\nclass MedianPool2d(nn.Module):\n    \"\"\" Median pool (usable as median filter when stride=1) module.\n    \n    Args:\n         kernel_size: size of pooling kernel, int or 2-tuple\n         stride: pool stride, int or 2-tuple\n         padding: pool padding, int or 4-tuple (l, r, t, b) as in pytorch F.pad\n         same: override padding and enforce same padding, boolean\n    \"\"\"\n    def __init__(self, kernel_size=3, stride=1, padding=0, same=False):\n        super(MedianPool2d, self).__init__()\n        self.k = _pair(kernel_size)\n        self.stride = _pair(stride)\n        self.padding = _quadruple(padding)  # convert to l, r, t, b\n        self.same = same\n\n    def _padding(self, x):\n        if self.same:\n            ih, iw = x.size()[2:]\n            if ih % self.stride[0] == 0:\n                ph = max(self.k[0] - self.stride[0], 0)\n            else:\n                ph = max(self.k[0] - (ih % self.stride[0]), 0)\n            if iw % self.stride[1] == 0:\n                pw = max(self.k[1] - self.stride[1], 0)\n            else:\n                pw = max(self.k[1] - (iw % self.stride[1]), 0)\n            pl = pw // 2\n            pr = pw - pl\n            pt = ph // 2\n            pb = ph - pt\n            padding = (pl, pr, pt, pb)\n        else:\n            padding = self.padding\n        return padding\n    \n    def forward(self, x):\n        # using existing pytorch functions and tensor ops so that we get autograd, \n        # would likely be more efficient to implement from scratch at C/Cuda level\n        x = F.pad(x, self._padding(x), mode='reflect')\n        x = x.unfold(2, self.k[0], self.stride[0]).unfold(3, self.k[1], self.stride[1])\n        x = x.contiguous().view(x.size()[:4] + (-1,)).median(dim=-1)[0]\n        return x\n","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"execution_failed":"2025-09-27T18:50:50.912Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile limb.py\n\n# https://www.kaggle.com/code/junkoda/limb-darkening\n\nimport torch\nimport math\n\nEPS = 1e-10\n\ndef area_torch(R, r, d):\n    d = torch.abs(d)\n\n    out = torch.zeros((len(d), len(R), len(r)), dtype=torch.float32, device=d.device)\n\n    d = d[:, :, None]\n    r = r[None, :, None]\n    R = R[None, None, :]\n\n    cos_alpha = (d ** 2 + R ** 2 - r ** 2) / (2 * R * d + EPS)\n    cos_beta = (r ** 2 + d ** 2 - R ** 2) / (2 * r * d + EPS)\n\n    cos_alpha = torch.clamp(cos_alpha, -1 + EPS, 1 - EPS)\n    cos_beta = torch.clamp(cos_beta, -1 + EPS, 1 - EPS)\n\n    overlap_sq = 2 * (R ** 2 * r ** 2 + r ** 2 * d ** 2 + d ** 2 * R ** 2) - (R ** 4 + r ** 4 + d ** 4)\n    overlap_sq = torch.clamp(overlap_sq, 0 + EPS, None)\n\n    out = R ** 2 * torch.acos(cos_alpha) + r ** 2 * torch.acos(cos_beta) - 0.5 * torch.sqrt(overlap_sq)\n\n    return out\n\n\ndef intensity_torch(d, c):\n    norm = (0.5\n            - c[:, 0]/10       # 1/2 - 1/(1/2+2) = 1/10\n            - c[:, 1]/6        # 1/6\n            - 3*c[:, 2]/14     # 3/14\n            - c[:, 3]/4        # 1/4\n            - 5*c[:, 4]/18     # 5/18 for p=2.5\n            - 3*c[:, 5]/10)    # 3/10 for p=3\n    norm = norm * 2 * torch.pi\n\n    d = torch.clamp(d, max=0.99995)\n    sqrtmu = (1 - d**2).clamp_min(0)**0.25  # μ^{1/2}\n\n    sqrtmu = sqrtmu[:, None]\n\n    num = (1\n           - c[None, :, 0]*(1 - sqrtmu)\n           - c[None, :, 1]*(1 - sqrtmu**2)\n           - c[None, :, 2]*(1 - sqrtmu**3)\n           - c[None, :, 3]*(1 - sqrtmu**4)\n           - c[None, :, 4]*(1 - sqrtmu**5)  # μ^{2.5}\n           - c[None, :, 5]*(1 - sqrtmu**6)) # μ^{3}\n\n    return num / norm[None, :]\n\n\ndef limb_darkening_torch(d, rp, c, *, nstep=41):\n    r_min = 0.0\n    r_max = 1.0\n\n    r = torch.linspace(r_min, r_max, nstep, dtype=d.dtype, device=d.device)\n    dr = torch.diff(r)\n    r_mid = r[:-1] + dr# * 0.5\n\n    A = area_torch(r, rp, d)\n\n    dA = torch.diff(A)\n    I = intensity_torch(r_mid, c)\n\n    I = I.permute(1, 0)\n\n    integ = torch.sum(I[None, :] * dA, dim=2)\n\n    return 1 - integ","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"execution_failed":"2025-09-27T18:50:50.913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile score.py\n\nimport numpy as np\nimport pandas as pd\nimport pandas.api.types\nimport scipy.stats\n\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\ndef competition_score(\n    solution: pd.DataFrame,\n    submission: pd.DataFrame,\n    row_id_column_name: str,\n    naive_mean: float,\n    naive_sigma: float,\n    fsg_sigma_true: float = 1e-6,\n    airs_sigma_true: float = 1e-5,\n    fgs_weight: float = 1,\n):\n    \"\"\"\n    This is a Gaussian Log Likelihood based metric. For a submission, which contains the predicted mean (x_hat) and variance (x_hat_std),\n    we calculate the Gaussian Log-likelihood (GLL) value to the provided ground truth (x). We treat each pair of x_hat,\n    x_hat_std as a 1D gaussian, meaning there will be 283 1D gaussian distributions, hence 283 values for each test spectrum,\n    the GLL value for one spectrum is the sum of all of them.\n\n    Inputs:\n        - solution: Ground Truth spectra (from test set)\n            - shape: (nsamples, n_wavelengths)\n        - submission: Predicted spectra and errors (from participants)\n            - shape: (nsamples, n_wavelengths*2)\n        naive_mean: (float) mean from the train set.\n        naive_sigma: (float) standard deviation from the train set.\n        fsg_sigma_true: (float) standard deviation from the FSG1 instrument for the test set.\n        airs_sigma_true: (float) standard deviation from the AIRS instrument for the test set.\n        fgs_weight: (float) relative weight of the fgs channel\n    \"\"\"\n\n    del solution[row_id_column_name]\n    del submission[row_id_column_name]\n\n    if submission.min().min() < 0:\n        raise ParticipantVisibleError('Negative values in the submission')\n    for col in submission.columns:\n        if not pandas.api.types.is_numeric_dtype(submission[col]):\n            raise ParticipantVisibleError(f'Submission column {col} must be a number')\n\n    n_wavelengths = len(solution.columns)\n    if len(submission.columns) != n_wavelengths * 2:\n        raise ParticipantVisibleError('Wrong number of columns in the submission')\n\n    y_pred = submission.iloc[:, :n_wavelengths].values\n    # Set a non-zero minimum sigma pred to prevent division by zero errors.\n    sigma_pred = np.clip(submission.iloc[:, n_wavelengths:].values, a_min=10**-15, a_max=None)\n    sigma_true = np.append(\n        np.array(\n            [\n                fsg_sigma_true,\n            ]\n        ),\n        np.ones(n_wavelengths - 1) * airs_sigma_true,\n    )\n    y_true = solution.values\n\n    GLL_pred = scipy.stats.norm.logpdf(y_true, loc=y_pred, scale=sigma_pred)\n    GLL_true = scipy.stats.norm.logpdf(y_true, loc=y_true, scale=sigma_true * np.ones_like(y_true))\n    GLL_mean = scipy.stats.norm.logpdf(y_true, loc=naive_mean * np.ones_like(y_true), scale=naive_sigma * np.ones_like(y_true))\n\n    # normalise the score, right now it becomes a matrix instead of a scalar.\n    ind_scores = (GLL_pred - GLL_mean) / (GLL_true - GLL_mean)\n\n    # print(GLL_mean.shape)\n\n    # print(GLL_pred, GLL_true)\n\n    # ind_scores[:, 0] = 0.5\n\n    weights = np.append(np.array([fgs_weight]), np.ones(len(solution.columns) - 1))\n    weights = weights * np.ones_like(ind_scores)\n    submit_score = np.average(ind_scores, weights=weights)\n    return float(np.clip(submit_score, 0.0, 1.0)), ind_scores, np.average(ind_scores, weights=np.append(np.array([fgs_weight]), np.ones(len(solution.columns) - 1)), axis=1)","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"execution_failed":"2025-09-27T18:50:50.913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile gauss_newton.py\n\nimport warnings\nfrom typing import Tuple, List, Dict, Optional, Callable\nimport torch\nfrom torch import nn\nfrom typing import Tuple, List\nimport torch\nfrom torch.func import vmap, jvp\nfrom functorch import make_functional_with_buffers\n\ndef _tree_add(a: Tuple[torch.Tensor, ...], b: Tuple[torch.Tensor, ...], alpha: float = 1.0):\n    return tuple(x + alpha * y for x, y in zip(a, b))\n\ndef _tree_scale(a: Tuple[torch.Tensor, ...], alpha: float):\n    return tuple(alpha * x for x in a)\n\ndef _tree_copy(a: Tuple[torch.Tensor, ...]):\n    return tuple(x.clone() for x in a)\n\ndef _tree_dot(a: Tuple[torch.Tensor, ...], b: Tuple[torch.Tensor, ...]) -> torch.Tensor:\n    return torch.stack([(x * y).reshape(-1).sum() for x, y in zip(a, b)]).sum()\n\ndef _tree_norm(a: Tuple[torch.Tensor, ...]) -> torch.Tensor:\n    return torch.sqrt(_tree_dot(a, a) + 1e-32)\n\ndef _flat(tup):\n    return torch.cat([t.reshape(-1) for t in tup])\n\ndef _unflat(vec, like_tup):\n    out = []\n    i = 0\n    for t in like_tup:\n        n = t.numel()\n        out.append(vec[i:i+n].view_as(t))\n        i += n\n    return tuple(out)\n\ndef residual_full_flat_weighted(\n    data: torch.Tensor,\n    weights: torch.Tensor,\n):\n    sqrt_w = torch.sqrt(weights)\n\n    def from_pred(pred: torch.Tensor) -> torch.Tensor:\n        residuals = (sqrt_w * (pred - data)).reshape(-1)\n        return residuals\n\n    return from_pred\n\ndef _make_basis(theta: Tuple[torch.Tensor, ...]) -> Tuple[Tuple[torch.Tensor, ...], int]:\n    flats = [t.reshape(-1) for t in theta]\n    sizes = [f.numel() for f in flats]\n    P = sum(sizes)\n    outs = []\n    offset = 0\n    for t, n_i in zip(theta, sizes):\n        B = torch.zeros((P,) + t.shape, device=t.device, dtype=t.dtype)\n        # fill diagonal block for this parameter\n        # reshape to [P, n_i] to place identity\n        Br = B.reshape(P, -1)\n        rows = torch.arange(offset, offset + n_i, device=t.device)\n        Br[rows, torch.arange(n_i, device=t.device)] = 1.0\n        outs.append(B)\n        offset += n_i\n    return tuple(outs), P\n\nclass GaussNewton:\n    def __init__(\n        self,\n        model: nn.Module,\n        residual_from_pred: Callable[[torch.Tensor], torch.Tensor],\n        lr: float = 1.0,\n        damping: float = 1e-3,\n        cg_max_iter: int = 50,\n        cg_tol: float = 1e-6,\n    ):\n        self.model = model\n        self.residual_from_pred = residual_from_pred\n        self.lr = lr\n        self.lam = damping\n        self.cg_max_iter = cg_max_iter\n        self.cg_tol = cg_tol\n\n        # functionalize the module once\n        with warnings.catch_warnings():\n            warnings.filterwarnings(\"ignore\", \".*make_functional_with_buffers.*deprecated.*\")\n            self.fmodel, self._params_all, self._buffers = make_functional_with_buffers(model)\n\n        # mask/index trainable leaves\n        self._mask: List[bool] = [p.requires_grad for p in model.parameters()]\n        self._train_idx: List[int] = [i for i, m in enumerate(self._mask) if m]\n\n        # current optimizer state: theta = tuple of trainable tensors\n        self.theta: Tuple[torch.Tensor, ...] = tuple(\n            self._params_all[i].detach().clone() for i in self._train_idx\n        )\n\n        # for copying updates back into the live module\n        self._live_params: List[nn.Parameter] = [p for p in model.parameters()]\n\n    def _merge_params(self, theta_tuple: Tuple[torch.Tensor, ...]) -> Tuple[torch.Tensor, ...]:\n        merged = list(self._params_all)\n        if len(theta_tuple) != len(self._train_idx):\n            raise RuntimeError(\n                f\"[GN] Trainable-count mismatch: got {len(theta_tuple)} values \"\n                f\"but expected {len(self._train_idx)} (indices {self._train_idx}).\"\n            )\n        for k, i in enumerate(self._train_idx):\n            merged[i] = theta_tuple[k]\n        return tuple(merged)\n\n    def _pred_at(self, theta_tuple: Tuple[torch.Tensor, ...]) -> torch.Tensor:\n        params_full = self._merge_params(theta_tuple)\n        out = self.fmodel(params_full, self._buffers)  # [1, T, S]\n        return out.squeeze(0)\n\n    def _residual(self, *theta_flat: torch.Tensor) -> torch.Tensor:\n        theta_tuple = tuple(theta_flat)\n        pred = self._pred_at(theta_tuple)\n        return self.residual_from_pred(pred)  # 1D vector\n\n    @torch.no_grad()\n    def _write_theta_into_model(self, theta_tuple: Tuple[torch.Tensor, ...]):\n        t = 0\n        for p, req in zip(self._live_params, self._mask):\n            if req:\n                p.copy_(theta_tuple[t])\n                t += 1\n        self._params_all = self._merge_params(theta_tuple)\n    \n    @torch.no_grad()\n    def step_dense_batched(self, chunk_cols: int = 0) -> dict:\n        theta = self.theta\n        r = self._residual(*theta)\n        r2 = float((r * r).sum().item())\n\n        basis, P = _make_basis(theta)\n\n        device = r.device\n        dtype = r.dtype\n        I = torch.eye(P, device=device, dtype=dtype)\n\n        G = torch.zeros((P, P), device=device, dtype=dtype)\n        g_vec = torch.zeros((P,), device=device, dtype=dtype)\n\n        if chunk_cols and chunk_cols < P:\n            start = 0\n            while start < P:\n                end = min(start + chunk_cols, P)\n                v_chunk = tuple(B[start:end] for B in basis)  # each has shape [B] + t.shape\n\n                Jv_block = vmap(lambda *v: jvp(self._residual, theta, v)[1],\n                                in_dims=(0,) * len(theta))(*v_chunk)  # [B, m]\n\n                G[start:end, start:end] += Jv_block @ Jv_block.T\n                g_vec[start:end] += Jv_block @ r\n\n                if start == 0:\n                    cache = [(start, end, Jv_block)]\n                else:\n                    for (s0, e0, Jv_prev) in cache:\n                        G[s0:e0, start:end] += Jv_prev @ Jv_block.T\n                        G[start:end, s0:e0] += (Jv_prev @ Jv_block.T).T\n                    cache.append((start, end, Jv_block))\n\n                start = end\n        else:\n            Jv_all = vmap(lambda *v: jvp(self._residual, theta, v)[1],\n                          in_dims=(0,) * len(theta))(*basis)\n            G = Jv_all @ Jv_all.T\n            g_vec = Jv_all @ r\n\n        lam = max(self.lam, 1e-8)\n        Gd = G + lam * I\n        L, info = torch.linalg.cholesky_ex(Gd)\n        if int(info) != 0:\n            for _ in range(5):\n                lam *= 10.0\n                Gd = G + lam * I\n                L, info = torch.linalg.cholesky_ex(Gd)\n                if int(info) == 0:\n                    break\n            if int(info) != 0:\n                delta_vec = torch.linalg.lstsq(Gd, -g_vec).solution\n            else:\n                delta_vec = torch.cholesky_solve((-g_vec).unsqueeze(1), L).squeeze(1)\n        else:\n            delta_vec = torch.cholesky_solve((-g_vec).unsqueeze(1), L).squeeze(1)\n\n        delta = _unflat(delta_vec, theta)\n\n        if chunk_cols and chunk_cols < P:\n            Jdelta = torch.zeros_like(r)\n            start = 0\n            while start < P:\n                end = min(start + (chunk_cols or P), P)\n                v_chunk = tuple(B[start:end] for B in basis)\n                Jv_block = vmap(lambda *v: jvp(self._residual, theta, v)[1],\n                                in_dims=(0,) * len(theta))(*v_chunk)  # [B, m]\n                delta_block = delta_vec[start:end]                    # [B]\n                Jdelta += (Jv_block.T @ delta_block)\n                start = end\n        else:\n            if 'Jv_all' not in locals():\n                Jv_all = vmap(lambda *v: jvp(self._residual, theta, v)[1],\n                              in_dims=(0,) * len(theta))(*basis)\n            Jdelta = Jv_all.T @ delta_vec                           # [m]\n\n        pred = 0.5 * (r2 - float(((r + Jdelta) ** 2).sum().item()))\n\n        theta_trial = tuple(ti + self.lr * di for ti, di in zip(theta, delta))\n        r_trial = self._residual(*theta_trial)\n        act = 0.5 * (r2 - float((r_trial * r_trial).sum().item()))\n        rho = act / (pred + 1e-32)\n\n        accepted = (rho > 1e-3) and (act > 0.0)\n        if accepted:\n            self.theta = theta_trial\n            self._write_theta_into_model(self.theta)\n            if rho > 0.75:\n                self.lam = max(self.lam * 0.3, 1e-8)\n            elif rho < 0.25:\n                self.lam = self.lam * 2.0\n        else:\n            self.lam = self.lam * 10.0\n\n        return {\n            \"accepted\": bool(accepted),\n            \"rho\": float(rho),\n            \"pred_red\": float(pred),\n            \"act_red\": float(act),\n            \"res_norm\": float((r_trial if accepted else r).norm().item()),\n            \"lam\": float(self.lam),\n        }    ","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"execution_failed":"2025-09-27T18:50:50.913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile model_new.py\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom utils import TransitTimings, get_ingress_egress_old, set_grad\nimport os\nimport copy\nimport tqdm\nimport limb\nfrom gauss_newton import GaussNewton, residual_full_flat_weighted\nimport torchvision\n\n\nclass ModelNewWrapper(nn.Module):\n    def __init__(\n            self,\n            star_info,\n            datas,\n            lr,\n            optimize_Rp=True,\n            gt_Rp_mean=None,\n            set_gt_Rp_mean=False,\n            gt_Rp_variation=None,\n            verbose=True,\n            optimize_p=True,\n            p=0.5,\n        ):\n        super().__init__()\n\n        self.star_info = star_info\n        self.datas = nn.Parameter(datas, requires_grad=False)\n        self.n_observations = datas.shape[0]\n        self.n_time = datas.shape[1]\n        self.n_spectra = datas.shape[2]\n        self.verbose = verbose\n\n        self.optimize_p = optimize_p\n        self.p = p\n\n        self.inouts = []\n        self.bin = 5625 // self.n_time\n        for data in self.datas:\n            in1, in2, in3, out1, out2, out3 = get_ingress_egress_old(data.cpu().numpy(), sigma=40)\n            inout = TransitTimings(\n                max(in1 - 100//self.bin, 0),\n                None,\n                None,\n                None,\n                None,\n                min(out3 + 100//self.bin, self.n_time),                \n            )\n            self.inouts.append(inout)\n\n        self.lr = lr\n        self.optimize_Rp = optimize_Rp\n        self.gt_Rp_mean = gt_Rp_mean\n        self.set_gt_Rp_mean = set_gt_Rp_mean\n        self.gt_Rp_variation = gt_Rp_variation\n        self.is_fit = False\n\n        self.model = ModelNew(\n            star_info=star_info,\n            n_time=self.n_time,\n            n_spectra=self.n_spectra,\n            inouts=self.inouts,\n            n_observations=self.n_observations,\n            n_rings=41 if self.n_spectra > 1 else 161,\n        )\n\n        star_spectrums_raw = []\n        for i in range(self.n_observations):\n            inout = self.inouts[i]\n            star_spectrum_raw = self.datas[i, :].mean(dim=0)\n            star_spectrums_raw.append(star_spectrum_raw)\n        self.star_spectrum_raw = torch.stack(star_spectrums_raw, dim=0).mean(dim=0)\n\n        self.initial_state_dict = copy.deepcopy(self.model.state_dict())\n\n    def load_from_cache(self, key, device, cache_path):\n        if os.path.exists(f\"{cache_path}/{key}.pt\"):\n            with open(f\"{cache_path}/{key}.pt\", 'rb') as f:\n                state_dict = torch.load(f, map_location=device)\n            self.model.load_state_dict(state_dict)\n            self.is_fit = True\n\n    def save_to_cache(self, key, cache_path):\n        os.makedirs(cache_path, exist_ok=True)\n        with open(f\"{cache_path}/{key}.pt\", 'wb') as f:\n            state_dict = self.model.state_dict()\n            torch.save(state_dict, f)\n\n    @torch.no_grad()\n    def get_loss(self):\n        residuals_gt = self.datas.reshape(-1)\n        a = torch.zeros_like(self.datas, device=self.datas.device)\n        a.fill_(1.0)\n        b = self.datas\n        if self.n_spectra > 1:\n            b = torchvision.transforms.GaussianBlur(kernel_size=(25, 1), sigma=(11, 11))(b[0, None, :, :])[None, 0, :, :]\n        weights = (a / b.std(dim=1, keepdim=True).pow(2)).expand(self.n_observations, self.n_time, self.n_spectra).reshape(-1)\n\n        model_pred = self.model.forward()\n        cur_loss = ((residuals_gt - model_pred).square() * weights).reshape(self.n_time, self.n_spectra)\n\n        return cur_loss\n\n    def fit_from_params(self, p, limb_coefs, star_spectrum_raw):\n        self.gt_maes = []\n        self.Rp_means = []\n        self.losses = []\n\n        self.model.load_state_dict(self.initial_state_dict)\n        self.model.train()\n        self.model.set_p(p)\n        self.model.set_coefs(limb_coefs)\n        self.model.star_spectrum_raw.data.copy_(star_spectrum_raw)\n        if self.set_gt_Rp_mean:\n            self.model.set_Rp_mean(self.gt_Rp_mean)\n        else:\n            in1 = self.inouts[0].in1\n            out3 = self.inouts[0].out3\n            Rp_rough = (1 - self.datas[0, in1:out3].mean(dim=0) / self.model.star_spectrum_raw).clip(min=0) ** 0.5\n            self.model.set_Rp_mean(Rp_rough.mean().item())\n        if self.gt_Rp_variation is not None:\n            self.model.Rp_variation.data.copy_(torch.tensor(self.gt_Rp_variation))\n\n        set_grad(False,\n            self.model.time_polynomial_coeffs,\n            self.model.Rp_variation,\n        )\n        set_grad(True,\n            self.model.limb_coefs,\n        )\n        set_grad(self.optimize_Rp,\n            self.model.Rp_mean,\n        )\n        set_grad(self.optimize_p, self.model.p)\n        \n        residuals_gt = self.datas.reshape(-1)\n        a = torch.zeros_like(self.datas, device=self.datas.device)\n        a.fill_(1.0)\n        b = self.datas\n        if self.n_spectra > 1:\n            b = torchvision.transforms.GaussianBlur(kernel_size=(25, 1), sigma=(11, 11))(b[0, None, :, :])[None, 0, :, :]\n        weights = (a / b.std(dim=1, keepdim=True).pow(2)).expand(self.n_observations, self.n_time, self.n_spectra).reshape(-1)\n        residual_fn = residual_full_flat_weighted(residuals_gt, weights)\n\n        lr = 1.0\n        cg_iter = 50\n\n        gn_args = dict(\n            model=self.model,\n            residual_from_pred=residual_fn,\n            lr=lr,\n            damping=1e+7,\n            cg_max_iter=cg_iter,\n            cg_tol=1e-6,\n        )\n        gn = GaussNewton(**gn_args)\n\n        poly_start = 70\n        n_iter = 220\n\n        if self.n_spectra == 1:\n            poly_start = 70\n            n_iter = 150\n\n        best_loss = float('inf')\n        best_state_dict = copy.deepcopy(self.model.state_dict())\n\n        for it in (bar := tqdm.trange(n_iter, disable=not self.verbose)):\n            if it == poly_start:\n                self.model.load_state_dict(best_state_dict)\n                set_grad(True, \n                    self.model.time_polynomial_coeffs,\n                )\n                gn = GaussNewton(**gn_args)\n            elif it == 180 and self.optimize_Rp:\n                if self.n_spectra == 1: break\n\n                self.model.load_state_dict(best_state_dict)\n                set_grad(False, self.model.Rp_mean)\n                set_grad(True, self.model.Rp_variation)\n                gn = GaussNewton(**gn_args)\n            elif gn.lam > 1e9:\n                continue\n\n            info = gn.step_dense_batched()\n\n            with torch.no_grad():\n                model_pred = self.model.forward()\n                cur_loss = ((residuals_gt - model_pred).square() * weights).mean()\n                self.Rp_means.append(self.model.get_Rp_mean().detach().cpu().numpy())\n                self.losses.append(cur_loss.item())\n\n                if cur_loss.item() < best_loss - 1e-12 and gn.lam < 1e10:\n                    best_loss = cur_loss.clone()\n                    best_state_dict = copy.deepcopy(self.model.state_dict())\n\n            if self.gt_Rp_mean is not None:\n                mae = (self.model.get_Rp() - self.gt_Rp_mean).abs().mean().item()# / self.gt_Rp_mean.item()\n                self.gt_maes.append(mae)\n                bar.set_description(f\"Loss: {cur_loss.item():.8f}; p: {self.model.get_p().item():.2f}; λ: {gn.lam:.1e}; mae: {mae:.7f}\")\n            else:\n                bar.set_description(f\"Loss: {cur_loss.item():.8f}; p: {self.model.get_p().item():.3f}; λ: {gn.lam:.1e}\")\n\n        self.gt_maes = np.array(self.gt_maes)\n        self.Rp_means = np.array(self.Rp_means)\n        self.losses = np.array(self.losses)\n\n        return best_loss if type(best_loss) == float else best_loss.item(), best_state_dict\n\n\n    def fit(self):\n        if self.is_fit: return\n\n        best_loss = float('inf')\n        best_state_dict = copy.deepcopy(self.model.state_dict())\n\n        params = [\n            {\"p\": self.p, \"limb_coefs\": 1, \"star_spectrum_raw\": self.star_spectrum_raw},\n        ]\n        \n        for param in params:\n            loss, state_dict = self.fit_from_params(**param)\n            if loss < best_loss:\n                best_loss = loss\n                best_state_dict = state_dict\n            torch.cuda.empty_cache()\n\n        if self.verbose: print(f\"Best loss: {best_loss:.6f}; p: {self.model.get_p().item():.3f}\")\n\n        self.model.load_state_dict(best_state_dict)\n        self.is_fit = True\n\n\ndef fit_line(y, x=None):\n    *batch, N = y.shape\n    if x is None:\n        x = torch.linspace(0.0, 1.0, N, dtype=y.dtype, device=y.device)\n    X = torch.stack([x, torch.ones_like(x)], dim=-1)             # (N, 2)\n\n    # Solve min_{a,b} ||Xa - y|| with a differentiable pseudo-inverse.\n    # Works for multiple series (batched) by treating them as multiple RHS.\n    Y = y.reshape(-1, N).mT                                       # (N, B)\n    beta = torch.linalg.pinv(X) @ Y                              # (2, B)\n    trend = (X @ beta).T.reshape(*batch, N)                      # (..., N)\n\n    a = beta[0].reshape(*batch)                                # slope\n    b = beta[1].reshape(*batch)                                # intercept\n    return a, b, trend\n\n\ndef detrend_linear(y, x=None, keep_level=\"remove_slope_keep_level\"):\n    a, b, trend = fit_line(y, x)\n    if keep_level == \"full_remove\":\n        y_dt = y - trend\n    elif keep_level == \"keep_mean\":\n        y_dt = y - trend + y.mean(dim=-1, keepdim=True)\n    elif keep_level == \"remove_slope_keep_level\":\n        y_dt = y - (trend - trend.mean(dim=-1, keepdim=True))\n    else:\n        raise ValueError(\"Unknown keep_level option.\")\n    return y_dt, (a, b, trend)\n\n\nclass ModelNew(nn.Module):\n    def __init__(\n            self,\n            star_info,\n            n_time,\n            n_spectra,\n            inouts,\n            limb_type=limb.intensity_torch,\n            n_observations=1,\n            degree_t=4,\n            degree_wl=2,\n            n_rings=41,\n        ):\n        super().__init__()\n\n        self.star_info = star_info\n        self.degree_t = degree_t\n        self.degree_wl = degree_wl\n        self.n_time = n_time\n        self.n_spectra = n_spectra\n        self.limb_type = limb_type\n        self.n_observations = n_observations\n        self.enable_variation = True\n        self.n_rings = n_rings\n\n        self.star_spectrum_raw = nn.Parameter(torch.zeros((n_spectra,), dtype=torch.float32), requires_grad=False)\n        # self.star_scaler = nn.Parameter(torch.tensor(1.0, dtype=torch.float32))\n        self.star_scaler = nn.Parameter(torch.ones(self.n_observations, 1, 1, dtype=torch.float32))\n\n        self.time_polynomial_coeffs = nn.Parameter(torch.zeros(self.n_observations, self.degree_t + 1, self.degree_wl + 1, dtype=torch.float32))\n        self.time_shift = nn.Parameter(torch.tensor(0.0, dtype=torch.float32), requires_grad=True)\n        self.time_scale = nn.Parameter(torch.tensor(1.0, dtype=torch.float32), requires_grad=True)\n        self.wl_shift = nn.Parameter(torch.tensor(0.0, dtype=torch.float32), requires_grad=True)\n        self.wl_scale = nn.Parameter(torch.tensor(1.0, dtype=torch.float32), requires_grad=True)\n\n        self.limb_coefs = nn.Parameter(torch.zeros(1, 6, dtype=torch.float32).expand(1, 6), requires_grad=False)\n\n        self.Rp_variation = nn.Parameter(torch.full((15,), 1, dtype=torch.float32), requires_grad=False)\n        self.Rp_mean = nn.Parameter(torch.tensor(1.0, dtype=torch.float32), requires_grad=False)\n\n        self.p = nn.Parameter(torch.tensor(0.0), requires_grad=False)\n\n        self.rmin = nn.Parameter(torch.tensor(0.0), requires_grad=True)\n        self.rmax = nn.Parameter(torch.tensor(0.0), requires_grad=True)\n\n        self.inouts = copy.deepcopy(inouts)\n        lengths = [inout.out3 - inout.in1 for inout in inouts]\n        self.transit_region_length = max(lengths)\n        for i, cur in enumerate(self.inouts):\n            length = cur.out3 - cur.in1\n            in1 = max(cur.in1 - (self.transit_region_length - length) // 2, 0)\n            out3 = min(in1 + self.transit_region_length, self.n_time)\n            self.inouts[i] = TransitTimings(in1, None, None, None, None, out3)\n\n        self.is_fit = False\n\n    def get_time_polynomial(self, coords):\n        time_steps = torch.arange(self.n_time, dtype=self.p.dtype, device=self.Rp_mean.device) / 10000\n        time_steps = (time_steps - self.time_shift) / self.time_scale\n    \n        wl_steps = torch.arange(self.n_spectra, dtype=self.p.dtype, device=self.Rp_mean.device) / 1000\n        wl_steps = (wl_steps - self.wl_shift) / self.wl_scale\n\n        Vx = torch.vander(time_steps, N=self.degree_t + 1, increasing=True)\n        Vy = torch.vander(wl_steps, N=self.degree_wl + 1, increasing=True)\n        poly = Vx[None] @ self.time_polynomial_coeffs @ Vy.T[None]\n\n        return poly\n\n    def get_coefs(self):\n        mask = torch.tensor([0.1, 0.1, 0.01, 0.01, 0.1, 0.1], device=self.limb_coefs.device)\n        return (self.limb_coefs * mask).square()\n    \n    def set_coefs(self, value):\n        value = torch.tensor(value, dtype=self.limb_coefs.dtype)\n        value = torch.sqrt(value)\n        self.limb_coefs.data.fill_(value)\n\n    def set_inout(self, inout):\n        self.inout = inout\n    \n    @torch.no_grad\n    def set_p(self, p):\n        p = torch.tensor(p, dtype=self.Rp_mean.dtype)\n        p = torch.log(p / (1 - p))\n        self.p.data.copy_(p / 4)\n\n    def get_p(self):\n        return F.sigmoid(self.p * 4)# * (1 - self.get_Rp().mean())\n\n    def get_dip_scale(self):\n        rmin = -(torch.sqrt((self.get_Rp().mean() + 1) ** 2 - self.get_p() ** 2) + self.rmin)[None].square()\n        rmax = +(torch.sqrt((self.get_Rp().mean() + 1) ** 2 - self.get_p() ** 2) + self.rmax)[None].square()\n\n        s = torch.linspace(0.0, 1.0, self.transit_region_length, device=self.Rp_mean.device)\n        s = s[:, None] * (rmax - rmin)[None, :] + rmin[None, :]\n\n        p = self.get_p()\n        d = torch.sqrt(p ** 2 + s ** 2)\n        l = limb.limb_darkening_torch(d, self.get_Rp(), self.get_coefs(), nstep=self.n_rings)\n\n        return l\n\n    def get_duration_star_info(self):\n        period = self.star_info[\"P\"]\n        sma = self.star_info[\"sma\"]\n        Rs = self.star_info[\"Rs\"]\n        RpRs = self.get_Rp().mean()\n        RpRs2 = self.get_depth().mean()\n        p = self.get_p()\n        duration = (period / (torch.pi * sma)) * torch.sqrt((1 + RpRs2) ** 2 - p ** 2)\n        return duration\n\n    def get_Rp(self, remove_trend=False):\n        mean = torch.ones(self.n_spectra, device=self.Rp_mean.device) * (self.Rp_mean / 10)\n        var = self.get_Rp_variation(remove_trend)\n        return var + mean    \n\n    def get_Rp_variation(self, remove_trend=False):\n        var = self.Rp_variation\n        var = F.interpolate(var[None, None, :], size=self.n_spectra, mode='linear', \n        align_corners=True).squeeze()\n        if remove_trend:\n            var, (a, b, trend) = detrend_linear(var.unsqueeze(0))\n            var = var.squeeze(0)\n        var = var - var.mean()\n        return var\n    \n    def set_Rp_mean(self, value):\n        self.Rp_mean.data.fill_(value * 10)\n\n    def get_Rp_mean(self):\n        return self.Rp_mean / 10\n\n    def get_depth(self, remove_trend=False):\n        return self.get_Rp(remove_trend) ** 2\n    \n    def forward(self):\n        pred = self.predictions()\n        return pred.reshape(-1)\n\n    def predictions(self):\n        time_steps = torch.arange(self.n_time, dtype=self.p.dtype, device=self.Rp_mean.device) / 10000\n        poly_t = self.get_time_polynomial(time_steps)\n        drift = 1 + poly_t\n\n        dip_scale = self.get_dip_scale()\n        dip_scale_extended = []\n        for i in range(self.n_observations):\n            in1 = self.inouts[i].in1\n            out3 = self.inouts[i].out3\n            dip_scale_extended.append(\n                torch.cat([\n                    torch.ones((in1, self.n_spectra), device=dip_scale.device),\n                    dip_scale,\n                    torch.ones((self.n_time - out3, self.n_spectra), device=dip_scale.device),\n                ], dim=0)\n            )\n        dip_scale_extended = torch.stack(dip_scale_extended, dim=0)\n\n        fact = dip_scale_extended * drift\n        star_scaler = 1 / fact.mean(dim=1, keepdim=True)\n\n        pred = self.star_spectrum_raw[None, None, :] * star_scaler * dip_scale_extended * drift\n    \n        return pred","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-27T18:50:50.913Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile utils.py\n\nimport numpy as np\nfrom scipy.ndimage import gaussian_filter1d\nfrom collections import namedtuple\nfrom scipy.signal import butter, sosfiltfilt\nfrom astropy.stats import sigma_clip\n\n\nTransitTimings = namedtuple(\"TransitTimings\", [\"in1\", \"in2\", \"in3\", \"out1\", \"out2\", \"out3\"])\n\n\ndef set_grad(requires_grad, *args):\n    for p in args:\n        p.requires_grad = requires_grad\n\n\ndef get_ingress_egress_old(data, sigma=40, safety=False):\n    data = outlier_sigma_clip(data, sigma=5)\n    data = data.mean(axis=1)\n    data_orig = data.copy()\n    data = gaussian_filter1d(data, sigma=sigma)\n\n    data = data[np.isfinite(data)]\n\n    d1 = np.diff(data, n=1)\n    d2 = np.diff(data, n=2) \n\n    if safety:\n        in2 = np.argmin(d1[:-100]) + 1  # +1 to account for the diff\n    else:\n        in2 = np.argmin(d1) + 1\n    out2 = np.argmax(d1) + 1  # +1 to account for the diff\n\n    in1 = np.argmin(d2[:in2]) + 1  # +1 to account for the diff\n    in3 = in2 + np.argmax(d2[in2:in2 + (out2 - in2) // 2]) + 1\n\n    out1 = out2 - (out2 - in2) // 2 + np.argmax(d2[out2 - (out2 - in2) // 2:out2]) + 1  # +1 to account for the diff\n    out3 = out2 + np.argmin(d2[out2:]) + 1\n\n    return TransitTimings(in1, in2, in3, out1, out2, out3)\n\n\ndef outlier_sigma_clip(data, sigma=4):\n    data = data.copy()\n\n    for i in range(data.shape[1]):\n        mask = sigma_clip(data[:, i], sigma=sigma).mask\n        data[mask, i] = np.nan\n        slce_new = data[:, i].copy()\n\n        for j in np.arange(len(mask))[mask]:\n            window = 2\n            while True:\n                if not np.all(np.isnan(\n                    data[max(0, j - window):min(len(mask), j + window + 1), i]\n                )):\n                    break\n                window += 1\n            \n            s = max(0, j - window)\n            f = min(len(mask), j + window + 1)\n            slce_new[j] = np.nanmean(data[s:f, i])\n\n        data[:, i] = slce_new\n\n    return data\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-27T18:50:50.913Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# https://www.kaggle.com/code/arsenypoyda/ariel-inference-9th-place\n\nfrom astropy.stats import sigma_clip\nimport itertools\n\nclass SignalCalibratorPublic:\n    def apply_linear_corr(self, 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            poly = np.poly1d(linear_corr[:, x, y])\n            clean_signal[:, x, y] = poly(clean_signal[:, x, y])\n        return clean_signal\n\n    def clean_dark(self, signal, dark, dt):\n        dark = np.tile(dark, (signal.shape[0], 1, 1))\n        signal -= dark* dt[:, np.newaxis, np.newaxis]\n        return signal\n    \n    def ADC_convert(self, signal, gain=0.4369, offset=-1000):\n        \"\"\"The Analog-to-Digital Conversion (adc) is performed by the detector to convert\n        the pixel voltage into an integer number. Since we are using the same conversion number \n        this year, we have simply hard-coded it inside. \"\"\"\n        signal = signal.astype(np.float64)\n        signal /= gain\n        signal += offset\n        return signal\n    \n    def calibrate_FGS(self, planet_id):\n        signal = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/FGS1_signal_0.parquet').to_numpy()\n        dark_frame = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/FGS1_calibration_0/dark.parquet').to_numpy()\n        dead_frame = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/FGS1_calibration_0/dead.parquet').to_numpy()\n        flat_frame = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/FGS1_calibration_0/flat.parquet').to_numpy()\n        linear_corr = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/FGS1_calibration_0/linear_corr.parquet').values.astype(np.float64).reshape((6, 32, 32))\n        signal = signal.reshape([-1, 32, 32]) \n\n        signal = self.ADC_convert(signal)\n        \n        dt = np.ones(len(signal))*0.1\n        dt[1::2] += 0.1    \n        \n        linear_corr_signal = self.apply_linear_corr(linear_corr, signal)\n        signal = self.clean_dark(linear_corr_signal, dark_frame, dt)  \n      \n        flat = flat_frame.reshape([1, 32, 32])\n        flat[dead_frame.reshape([1, 32, 32])] = np.nan\n        signal = signal / flat\n\n        signal = signal[:, 10:22 , 10:22]\n\n        mean_signal = np.nanmean(signal, axis=(1, 2))\n        cds_signal = (mean_signal[1::2] - mean_signal[0::2])\n\n        cds_signal = cds_signal.reshape(-1, 12).mean(axis=1)\n\n        np.save(f'{CAL_ROOT}/{planet_id}_FGS_0.npy', cds_signal)\n\n    def calibrate(self, planet_id):\n        cut_inf, cut_sup = 0, 356\n\n        signal = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/AIRS-CH0_signal_0.parquet').to_numpy()\n        dark_frame = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/AIRS-CH0_calibration_0/dark.parquet').to_numpy()\n        dead_frame = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/AIRS-CH0_calibration_0/dead.parquet').to_numpy()\n        flat_frame = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/AIRS-CH0_calibration_0/flat.parquet').to_numpy()\n        linear_corr = pd.read_parquet(f'{DATA_ROOT}/{SPLIT}/{planet_id}/AIRS-CH0_calibration_0/linear_corr.parquet').values.astype(np.float64).reshape((6, 32, 356))\n        signal = signal.reshape([11250, 32, 356]) \n\n        signal = self.ADC_convert(signal)\n        \n        hot = sigma_clip(\n            dark_frame, sigma=5, maxiters=5\n        ).mask\n        \n        signal = signal[:, :, cut_inf:cut_sup] \n        dt = np.ones(len(signal))*0.1 \n        dt[1::2] += 4.5 #@bilzard idea\n        linear_corr = linear_corr[:, :, cut_inf:cut_sup]\n        dark_frame = dark_frame[:, cut_inf:cut_sup]\n        dead_frame = dead_frame[:, cut_inf:cut_sup]\n        flat_frame = flat_frame[:, cut_inf:cut_sup]\n        hot = hot[:, cut_inf:cut_sup]\n            \n        linear_corr_signal = self.apply_linear_corr(linear_corr, signal)\n        signal = self.clean_dark(linear_corr_signal, dark_frame, dt)\n        \n        flat = flat_frame.reshape([1, 32, cut_sup-cut_inf])\n        flat[dead_frame.reshape([1, 32, cut_sup-cut_inf])] = np.nan\n        # flat[hot.reshape([1, 32, cut_sup-cut_inf])] = np.nan\n        signal = signal / flat\n\n        # signal = signal[:, 10:22 , :]\n        signal = signal[:, 8:24 , :]\n        mean_signal = np.nanmean(signal, axis=1)\n        cds_signal = (mean_signal[1::2] - mean_signal[0::2])\n\n        np.save(f'{CAL_ROOT}/{planet_id}_AIRS_0.npy', cds_signal)\n        print(\"Saved\", f'{CAL_ROOT}/{planet_id}_AIRS_0.npy')\n    \n    def calibrate_parallel(self, planet_ids):        \n        feats = joblib.Parallel(n_jobs=4, backend=\"loky\")(\n            joblib.delayed(self.calibrate)(planet_id) for planet_id in tqdm.tqdm(list(planet_ids))\n        )\n\n        feats = joblib.Parallel(n_jobs=4, backend=\"loky\")(\n            joblib.delayed(self.calibrate_FGS)(planet_id) for planet_id in tqdm.tqdm(list(planet_ids))\n        )","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-27T18:50:50.913Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport joblib\nimport tqdm\nimport os\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom colorama import init as colorama_init\nfrom colorama import Fore\nfrom colorama import Style\ncolorama_init()\n\nimport math\nfrom scipy.ndimage import gaussian_filter1d\n\nfrom model_new import ModelNewWrapper\nfrom utils import get_ingress_egress_old, TransitTimings, outlier_sigma_clip\n\n\nIS_KAGGLE = os.path.exists('/kaggle/input')\n\nMODE = \"normal\"\n# MODE = \"train\"\n# MODE = \"inference\"\n\nBIN = 1\nDEVICE = \"cuda\"\nCOORDS_SCALE = 10000\nPLOT = False\nVERBOSE = (MODE == 'normal')\n\nif not IS_KAGGLE:\n    DATA_ROOT = \"/home/me/brain/ariel/data_2025_new/\"\n    CAL_ROOT = \"/home/me/brain/ariel/data_2025_cal_new/\"\n    CACHE_PATH = \"/home/me/brain/ariel/cache_time/\"\nelse:\n    DATA_ROOT = \"/kaggle/input/ariel-data-challenge-2025/\"\n    CAL_ROOT = \"/kaggle/tmp/ariel-data-challenge-2025-cal/\"\n    CACHE_PATH = \"/kaggle/tmp/ariel-data-challenge-2025-cache/\"\n\nif not os.path.exists(CAL_ROOT):\n    os.makedirs(CAL_ROOT)\n\nif not os.path.exists(CACHE_PATH):\n    os.makedirs(CACHE_PATH)\n\nSPLIT = \"train\"\n\ndata_folder = f'{DATA_ROOT}/{SPLIT}'\nstar_info = pd.read_csv(f\"{DATA_ROOT}/{SPLIT}_star_info.csv\")\nstar_info[\"planet_id\"] = star_info[\"planet_id\"].astype(int)\naxis_info = pd.read_parquet(f'{DATA_ROOT}/axis_info.parquet')\nwavelengths = pd.read_csv(f'{DATA_ROOT}/wavelengths.csv')\n\nif SPLIT == \"train\":\n    train_labels = pd.read_csv(f\"{DATA_ROOT}/train.csv\", index_col='planet_id')\n\nangles = []\n\nstar_info[\"p\"] = np.cos(np.radians(star_info[\"i\"])) * star_info[\"sma\"]","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-27T18:50:50.914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"planet_ids = star_info[\"planet_id\"].tolist()[:10]","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-27T18:50:50.914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"garrus = SignalCalibratorPublic()\ngarrus.calibrate_parallel(planet_ids)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-27T18:50:50.914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mu_all = []\nsigma_all = []\ndata = []\n\nlogs_airs = []\nlogs_fgs = []\nmus = []\nsigmas = []\ncleans_airs = []\ncleans_fgs = []\n\nimport time\n\nimport pickle\nfrom sklearn.decomposition import PCA\nwith open(\"/kaggle/input/calibrated/pca.pkl\", \"rb\") as f:\n    pca = pickle.load(f)\n\nimport catboost\n\nmodel_dmu_airs = catboost.CatBoostRegressor()\nmodel_dmu_airs.load_model(\"/kaggle/input/calibrated/boost_dmu_airs_v3.cbm\")\n\nmodel_dsigma_airs = catboost.CatBoostRegressor()\nmodel_dsigma_airs.load_model(\"/kaggle/input/calibrated/boost_dsigma_airs_v3.cbm\")\n\nmodel_dmu_fgs = catboost.CatBoostRegressor()\nmodel_dmu_fgs.load_model(\"/kaggle/input/calibrated/boost_dmu_fgs_v3.cbm\")\n\nmodel_dsigma_fgs = catboost.CatBoostRegressor()\nmodel_dsigma_fgs.load_model(\"/kaggle/input/calibrated/boost_dsigma_fgs_v3.cbm\")\n\ndef get_pred(planet_id, sample_i):\n    device = DEVICE\n    start_time = time.time()\n\n    wl = 37\n    \n    # ==== READ AIRS ==== #\n\n    datas_airs = []\n    for i in [0]:\n        file = f\"{CAL_ROOT}/{planet_id}_AIRS_{i}.npy\"\n        if not os.path.exists(file): continue\n        data_airs = np.load(file)[:, wl:wl+282]\n        data_airs = outlier_sigma_clip(data_airs, sigma=5)\n        datas_airs.append(data_airs)\n    datas_airs = np.stack(datas_airs, axis=0)\n    datas_airs_orig = datas_airs.copy()\n    datas_airs = datas_airs.reshape(len(datas_airs), -1, 5, 282).mean(axis=-2)\n    datas_airs = datas_airs.reshape(1, datas_airs.shape[1], -1, 6).mean(axis=-1)\n\n    # ==== READ FGS ==== #\n\n    datas_fgs = []\n    for i in [0]:\n        file = f\"{CAL_ROOT}/{planet_id}_FGS_{i}.npy\"\n        if not os.path.exists(file): continue\n        data_fgs = np.load(file)[:, None]\n        data_fgs = outlier_sigma_clip(data_fgs, sigma=5)\n        datas_fgs.append(data_fgs)\n    datas_fgs = np.stack(datas_fgs, axis=0)\n\n    in1, in2, in3, out1, out2, out3 = get_ingress_egress_old(datas_airs_orig[0], sigma=40)\n\n    if VERBOSE: print(f\"{Fore.RED}({sample_i}/{len(planet_ids)}) Planet ID: {planet_id}{Style.RESET_ALL}\")\n\n    if SPLIT == \"train\":\n        gt = train_labels[train_labels.index == planet_id].values.squeeze(0)#[1:]\n        gt_airs = gt[1:]\n        gt_fgs = gt[:1]\n        gt = np.concatenate([gt_fgs, gt_airs])\n\n    # ==== FIT AIRS ==== #\n\n    model_airs = ModelNewWrapper(\n        star_info=star_info[star_info[\"planet_id\"] == planet_id].iloc[0],\n        datas=torch.tensor(datas_airs, dtype=torch.float32),\n        lr=1.0,\n        # gt_Rp_mean=gt_airs.mean() ** 0.5,\n        # optimize_Rp=False,\n        # set_gt_Rp_mean=True,\n        verbose=VERBOSE,\n    )\n    model_airs.double(); id = f\"{planet_id}_AIRS_double_bin\"\n    model_airs.to(device)\n    model_airs.fit()\n\n    mu_airs = model_airs.model.get_Rp_mean().expand(282).pow(2).detach().cpu().numpy()[::-1]\n    model_airs.model.n_spectra = 282\n    mu_airs_var = model_airs.model.get_depth(remove_trend=True).detach().cpu().numpy()[::-1]\n    model_airs.model.n_spectra = 47\n    mu_airs_var[-50:] = mu_airs_var[-50]\n\n    mu_airs_var_pca = pca.inverse_transform(pca.transform((mu_airs_var - mu_airs_var.mean())[None]))[0] + mu_airs_var.mean()\n    mu_airs_var_pca = gaussian_filter1d(mu_airs_var_pca, sigma=5)\n    \n    mu_airs_var = gaussian_filter1d(mu_airs_var, sigma=5)    \n\n    # ==== FIT FGS ==== #\n\n    model_fgs = ModelNewWrapper(\n        star_info=star_info[star_info[\"planet_id\"] == planet_id].iloc[0],\n        datas=torch.tensor(datas_fgs, dtype=torch.float32),\n        lr=1.0,\n        # gt_Rp_mean=gt_fgs.mean() ** 0.5,\n        optimize_p=False,\n        p=model_airs.model.get_p().item(),\n        # optimize_Rp=False,\n        # set_gt_Rp_mean=True,\n        verbose=VERBOSE,\n    )\n    model_fgs.double(); id = f\"{planet_id}_FGS_double_bin\"\n    model_fgs.to(device)\n    model_fgs.fit()\n\n    clean_fgs = True\n    mu_fgs = (model_fgs.model.get_Rp_mean().pow(2).detach().cpu().numpy()[None] * 0.6 + mu_airs_var[:1] * 0.4)\n    if (mu_fgs / mu_airs_var[:1]) > 1.3 or (mu_fgs / mu_airs_var[:1]) < 0.7:\n        clean_fgs = False\n        mu_fgs = mu_airs_var[:1]\n\n    # ==== COMBINE ==== #\n\n    mu = np.concatenate([mu_fgs, mu_airs_var_pca])\n    sigma = np.ones_like(mu) * (mu_airs_var.std() * 1.8 - mu_airs_var_pca.std() * 0.5)\n\n    clean = True\n    if model_airs.model.get_p().item() > 0.75:\n        clean = False\n        mu = mu / 1.05\n        sigma = sigma * 7.0\n\n    sigma[0] *= 1.8\n\n    if in1 < 150 and (datas_airs_orig.shape[1] - out3) < 150:\n        clean = False\n        # mu = mu * 0 + 0.01\n        sigma = sigma * 0 + 0.015\n\n    if np.isnan(mu).any() or np.isinf(mu).any() or sigma.mean() < 1e-5:\n        clean = False\n        print(\"NaN or Inf in mu for planet\", planet_id)\n        mu = np.ones(283) * 0.01\n        sigma = np.ones(283) * 0.01\n\n    old_sigma = sigma.copy()\n\n    features_airs = np.concatenate([\n        mu_airs_var.std()[None],\n        mu_airs_var_pca.std()[None],\n        model_airs.model.get_p()[None].detach().cpu().numpy(),\n        model_airs.model.get_coefs().cpu().detach().numpy().flatten(),\n        model_airs.model.time_polynomial_coeffs.detach().cpu().numpy().flatten(),\n        np.array(model_airs.model.transit_region_length)[None] / 1000,\n        mu_airs.mean()[None],\n        mu_airs_var,\n        star_info[star_info[\"planet_id\"] == planet_id].iloc[0].values[2:],\n        model_airs.get_loss().mean()[None].cpu().numpy(),\n        model_airs.get_loss().std()[None].cpu().numpy(),\n    ])\n\n    features_fgs = np.concatenate([\n        model_fgs.model.get_p()[None].detach().cpu().numpy(),\n        model_fgs.model.get_coefs().detach().cpu().numpy().flatten(),\n        model_fgs.model.time_polynomial_coeffs.detach().cpu().numpy().flatten(),\n        np.array(model_fgs.model.transit_region_length)[None] / 1000,\n        mu_fgs.mean()[None],\n        star_info[star_info[\"planet_id\"] == planet_id].iloc[0].values[2:],\n        model_fgs.get_loss().mean()[None].cpu().numpy(),\n        model_fgs.get_loss().std()[None].cpu().numpy(),\n    ])\n\n    if clean:\n        dsigma_airs = model_dsigma_airs.predict(features_airs[None, :])[0]\n        if (0.3 < (dsigma_airs + old_sigma.mean()) / old_sigma.mean()):\n            sigma[1:] = sigma[1:] + dsigma_airs\n\n        dmu_airs = model_dmu_airs.predict(features_airs[None, :])[0]\n        if (np.abs(dmu_airs) / mu_airs.mean() < 0.01):\n            mu[1:] = mu[1:] + dmu_airs\n            sigma[1:] = sigma[1:] * 0.9\n\n        if clean_fgs:\n            dsigma_fgs = model_dsigma_fgs.predict(features_fgs[None, :])[0]\n            if (0.7 < (dsigma_fgs + sigma[0]) / sigma[0]):\n                sigma[0] = sigma[0] + dsigma_fgs\n\n            dmu_fgs = model_dmu_fgs.predict(features_fgs[None, :])[0]\n            if (np.abs(dmu_fgs) / mu_fgs.mean() < 0.02):\n                mu[0] = mu[0] + dmu_fgs    \n\n    finish_time = time.time()\n    if VERBOSE: print(f\"{Fore.GREEN}Processed {planet_id} in {finish_time - start_time:.2f} seconds{Style.RESET_ALL}\")\n\n    del model_airs\n    del model_fgs\n    torch.cuda.empty_cache()\n\n    if np.isnan(mu).any() or np.isinf(mu).any() or sigma.mean() < 1e-5:\n        print(\"NaN or Inf in mu for planet\", planet_id, \"AGAIN\")\n        mu = np.ones(283) * 0.01\n        sigma = np.ones(283) * 0.01\n\n    return mu, sigma\n\nif MODE == \"inference\":\n    feats = joblib.Parallel(n_jobs=8, backend=\"loky\")(\n        joblib.delayed(get_pred)(planet_id, i) for i, planet_id in enumerate(tqdm.tqdm(list(planet_ids)))\n    )\n    mu_all, sigma_all = zip(*feats)\nelse:\n    for i, planet_id in enumerate(tqdm.tqdm(planet_ids, disable=VERBOSE)):\n        try:\n            mu, sigma = get_pred(planet_id, i)\n            mu_all.append(mu)\n            sigma_all.append(sigma)\n        except KeyboardInterrupt:\n            print(\"Interrupted by user\")\n            exit()\n\n\nmu_pred = np.stack(mu_all, axis=0)\nsigma_pred = np.stack(sigma_all, axis=0)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-27T18:50:50.914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from score import competition_score\n\ndef format_mu(planet_codes, mu, wavelengths):\n    return pd.DataFrame(mu.clip(0, None), index=planet_codes, columns=wavelengths.columns)\n\ndef format_sigma(planet_codes, sigma):\n    return pd.DataFrame(sigma, index=planet_codes, columns=[f\"sigma_{i}\" for i in range(1, 284)])\n\ndef format_to_submission(planet_codes, mus, sigmas):\n    mus = pd.concat(mus)\n    sigmas = pd.concat(sigmas)\n    df_index = pd.DataFrame({'planet_id': list(planet_codes)}, index=planet_codes)\n    df_submission = df_index.join(mus).join(sigmas)\n    return df_submission.reset_index(drop=True)\n\nmu_df = format_mu(planet_ids, mu_pred, wavelengths)\nsigma_df = format_sigma(planet_ids, sigma_pred)\n\nsubmission_df = format_to_submission(\n    star_info[star_info[\"planet_id\"].isin(planet_ids)][\"planet_id\"], \n    mus=[mu_df],\n    sigmas=[sigma_df],\n)\n\nif SPLIT == \"train\":\n    train_labels = train_labels[train_labels.index.isin(planet_ids)]\n\n    print(f\"MSE: {np.square(train_labels.values[:, 1:] - mu_df.loc[train_labels.index].values[:, 1:]).mean()}\")\n    score, ind_scores, ind_avg = competition_score(\n        train_labels.copy().reset_index(),\n        submission_df.copy(),\n        row_id_column_name=\"planet_id\",\n        naive_mean=train_labels.values.mean(),\n        naive_sigma=train_labels.values.std(),\n        fsg_sigma_true=1e-6,\n        airs_sigma_true=1e-5,\n        fgs_weight=57.846,\n        # fgs_weight=0,\n    )\n    print('SCORE:', score)\n\n\nsubmission_df.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-27T18:50:50.914Z"}},"outputs":[],"execution_count":null}]}