{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.13.5"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":101849,"databundleVersionId":13093295,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"2c420e48","cell_type":"markdown","source":"# NeurIPS Ariel Data Challenge 2025 — End-to-End Baseline\nThis notebook provides a CPU-friendly baseline pipeline to extract exoplanet spectra from Ariel simulated data and generate a valid `submission.csv` for the competition. It is designed to run on Kaggle with internet disabled.","metadata":{}},{"id":"94165bfe","cell_type":"code","source":"# 1) Set Up Environment and Paths\nimport os, sys, math, json, random, platform, pathlib, gc, time\nfrom pathlib import Path\nNOTEBOOK_NAME = \"NeurIPS_Ariel_2025_baseline\"\nprint(f\"Running: {NOTEBOOK_NAME}\")\n\n# Detect Kaggle environment and define input/output directories\nKAGGLE_INPUT = Path(\"/kaggle/input\")\nKAGGLE_WORKING = Path(\"/kaggle/working\")\nLOCAL_ROOT = Path(\".\")\nIN_KAGGLE = KAGGLE_INPUT.exists()\n\nif IN_KAGGLE:\n    INPUT_DIR = KAGGLE_INPUT\n    WORK_DIR = KAGGLE_WORKING\nelse:\n    INPUT_DIR = LOCAL_ROOT\n    WORK_DIR = Path(\"./_working\")\nWORK_DIR.mkdir(parents=True, exist_ok=True)\n\n# Auto-detect dataset root inside /kaggle/input if present\nif IN_KAGGLE:\n    detected = None\n    try:\n        for d in KAGGLE_INPUT.iterdir():\n            if d.is_dir() and (d / 'sample_submission.csv').exists():\n                detected = d\n                break\n        if detected is not None:\n            INPUT_DIR = detected\n            print(\"Detected dataset root:\", INPUT_DIR)\n        else:\n            print(\"Warning: Could not detect dataset folder; using /kaggle/input root.\")\n    except Exception as e:\n        print(\"Dataset detection error:\", e)\n\nSUBMISSION_FILENAME = \"submission.csv\"\nSUBMISSION_PATH = WORK_DIR / SUBMISSION_FILENAME\nCHECKPOINT_DIR = WORK_DIR / \"checkpoints\"\nCACHE_DIR = WORK_DIR / \"cache\"\nfor d in (CHECKPOINT_DIR, CACHE_DIR):\n    d.mkdir(parents=True, exist_ok=True)\n\n# Determinism and threads\nos.environ.setdefault(\"PYTHONHASHSEED\", \"42\")\nos.environ.setdefault(\"OMP_NUM_THREADS\", \"1\")\nos.environ.setdefault(\"MKL_NUM_THREADS\", \"1\")\nos.environ.setdefault(\"OPENBLAS_NUM_THREADS\", \"1\")\nos.environ.setdefault(\"NUMEXPR_NUM_THREADS\", \"1\")\nrandom.seed(42)\nprint(\"Env ready. Kaggle:\", IN_KAGGLE, \"Input:\", str(INPUT_DIR), \"Work:\", str(WORK_DIR))","metadata":{},"outputs":[],"execution_count":null},{"id":"9eaa2aae","cell_type":"code","source":"# 2) Import Libraries and Configure Precision/Seed\nimport numpy as np\nnp.set_printoptions(precision=6, suppress=True)\nnp.random.seed(42)\n\nimport pandas as pd\npd.options.display.max_rows = 50\npd.options.display.width = 140\n\ntry:\n    import pyarrow as pa\n    import pyarrow.parquet as pq\nexcept Exception as e:\n    pa = None; pq = None; print(\"PyArrow not available, will fall back to pandas read_parquet if possible.\")\n\nfrom dataclasses import dataclass\nfrom typing import Iterator, Optional, Tuple, Dict, Any, List\n\n# Plotting\nimport matplotlib.pyplot as plt\ntry:\n    import seaborn as sns\n    sns.set_context(\"notebook\")\n    sns.set_style(\"whitegrid\")\nexcept Exception:\n    sns = None\n\n# SciPy/Statsmodels/Sklearn optional imports\ntry:\n    from scipy import signal, optimize, stats, linalg\nexcept Exception as e:\n    signal = optimize = stats = linalg = None\ntry:\n    import statsmodels.api as sm\nexcept Exception as e:\n    sm = None\ntry:\n    from sklearn.linear_model import Ridge, BayesianRidge\n    from sklearn.model_selection import GroupKFold\nexcept Exception as e:\n    Ridge = BayesianRidge = GroupKFold = None\n\n# Optional acceleration\nUSE_NUMBA = False\ntry:\n    import numba\n    USE_NUMBA = True\nexcept Exception:\n    USE_NUMBA = False\nprint(\"Numba:\", USE_NUMBA)\n\nFLOAT = np.float64","metadata":{},"outputs":[],"execution_count":null},{"id":"aa3a7591","cell_type":"code","source":"# 3) Load Competition Metadata (CSV/Parquet)\ndef try_read_csv(path: Path) -> Optional[pd.DataFrame]:\n    f = path if path.exists() else None\n    if f is None:\n        return None\n    try:\n        return pd.read_csv(f)\n    except Exception as e:\n        print(\"Failed to read:\", f, e)\n        return None\n\ndef try_read_parquet(path: Path) -> Optional[pd.DataFrame]:\n    f = path if path.exists() else None\n    if f is None:\n        return None\n    try:\n        return pd.read_parquet(f)\n    except Exception as e:\n        print(\"Failed to read:\", f, e)\n        return None\n\nmeta = {}\nmeta['train'] = try_read_csv(INPUT_DIR / 'train.csv')\nmeta['wavelengths'] = try_read_csv(INPUT_DIR / 'wavelengths.csv')\n_axis_pq = try_read_parquet(INPUT_DIR / 'axis_info.parquet')\nmeta['axis_info'] = _axis_pq if _axis_pq is not None else try_read_csv(INPUT_DIR / 'axis_info.csv')\nmeta['adc_info'] = try_read_csv(INPUT_DIR / 'adc_info.csv')\nmeta['train_star'] = try_read_csv(INPUT_DIR / 'train_star_info.csv')\nmeta['test_star'] = try_read_csv(INPUT_DIR / 'test_star_info.csv')\nmeta['sample_submission'] = try_read_csv(INPUT_DIR / 'sample_submission.csv')\n\nfor k,v in meta.items():\n    if v is not None:\n        print(k, v.shape)\n    else:\n        print(k, None)\n\n# Basic validation (soft, do not assert to keep pipeline running)\nif meta['wavelengths'] is not None:\n    required_wl_cols = {'instrument','wavelength','index'}\n    if not required_wl_cols.issubset(set(meta['wavelengths'].columns)):\n        print(\"Warning: wavelengths.csv missing expected columns:\", required_wl_cols, \"found:\", list(meta['wavelengths'].columns))\nelse:\n    print(\"Warning: wavelengths.csv not found; proceeding without explicit wavelength grid.\")\n\nif meta['train'] is not None:\n    if 'planet_id' not in meta['train'].columns:\n        print(\"Warning: train.csv missing planet_id column; downstream CV disabled.\")\n\nif meta['sample_submission'] is not None:\n    required_sub_cols = {'planet_id','instrument','index','mu','sigma'}\n    if not required_sub_cols.issubset(set(meta['sample_submission'].columns)):\n        print(\"Warning: sample_submission.csv missing expected columns:\", required_sub_cols, \"found:\", list(meta['sample_submission'].columns))\nelse:\n    print(\"Warning: sample_submission.csv not found; will build submission skeleton from predictions only.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"c5cdc95b","cell_type":"code","source":"def compute_linear_coeffs(df_poly: pd.DataFrame, max_degree: int = 2) -> np.ndarray:\n    # Try to parse polynomial coefficients from calibration if available.\n    # Expected flexible formats:\n    # - Rows = pixels, columns named c0,c1,c2,...\n    # - Otherwise, fall back to identity (no correction)\n    if df_poly is not None:\n        cols = [c for c in df_poly.columns if isinstance(c, str) and c.startswith('c') and c[1:].isdigit()]\n        if cols:\n            # Sort columns by degree order c0,c1,c2...\n            cols_sorted = sorted(cols, key=lambda x: int(x[1:]))\n            coeffs = df_poly[cols_sorted].to_numpy(dtype=np.float64)\n            # Ensure up to max_degree+1 columns\n            if coeffs.shape[1] < (max_degree + 1):\n                # Pad with zeros\n                pad = np.zeros((coeffs.shape[0], (max_degree + 1) - coeffs.shape[1]), dtype=np.float64)\n                coeffs = np.hstack([coeffs, pad])\n            elif coeffs.shape[1] > (max_degree + 1):\n                coeffs = coeffs[:, : (max_degree + 1)]\n            # Sanity: if all-zero, use identity\n            if np.allclose(coeffs, 0.0):\n                n_pix = df_poly.shape[0]\n                coeffs = np.zeros((n_pix, max_degree+1), dtype=np.float64)\n                coeffs[:,1] = 1.0\n            return coeffs\n    # Fallback: identity transform y = x\n    n_pix = df_poly.shape[1] if (df_poly is not None and df_poly.shape[1] > 0) else 1024\n    coeffs = np.zeros((n_pix, max_degree+1), dtype=np.float64)\n    coeffs[:,1] = 1.0\n    return coeffs","metadata":{},"outputs":[],"execution_count":null},{"id":"9eaed711","cell_type":"code","source":"def airs_unflatten(arr: np.ndarray) -> np.ndarray:\n    # arr: (n_frames, 11392) -> (n_frames, 32, 356)\n    return arr.reshape(arr.shape[0], 32, 356)\n\ndef median_collapsed_image(frames_3d: np.ndarray) -> np.ndarray:\n    return np.nanmedian(frames_3d, axis=0)\n\ndef find_trace_center_y(med_img: np.ndarray) -> int:\n    # crude approach: sum across dispersion (x) and pick max row\n    prof = np.nansum(med_img, axis=1)  # shape (32,)\n    return int(np.nanargmax(prof))\n\ndef crop_airs_x(frames_3d: np.ndarray, x0: int = 39, x1: int = 321) -> np.ndarray:\n    # Crop dispersion axis to recommended region\n    x0 = max(0, int(x0)); x1 = min(frames_3d.shape[2], int(x1))\n    return frames_3d[:, :, x0:x1]\n\ndef optimal_extract_1d(frames_3d: np.ndarray, center_y: int, half_height: int = 3) -> np.ndarray:\n    # Simple box/optimal hybrid extraction, assumes frames_3d already cropped on x if desired\n    y0 = max(0, center_y - half_height)\n    y1 = min(frames_3d.shape[1], center_y + half_height + 1)\n    sub = frames_3d[:, y0:y1, :]  # (n, h, W)\n    # Profile per column\n    prof = np.nanmedian(sub, axis=0)  # (h, W)\n    prof = prof / (np.nanmax(prof, axis=0, keepdims=True) + 1e-12)\n    # Weighted sum\n    w = prof[None, :, :]\n    spec = np.nansum(sub * w, axis=1)  # (n, W)\n    return spec\n\ndef extract_airs_spectra_batch(calibrated_batch: np.ndarray) -> np.ndarray:\n    # Input (n_frames, 11392) -> output (n_frames, W_crop)\n    frames = airs_unflatten(calibrated_batch)\n    # Crop dispersion axis to reduce noise and edges\n    frames = crop_airs_x(frames, 39, 321)\n    med_img = median_collapsed_image(frames)\n    cy = find_trace_center_y(med_img)\n    spec = optimal_extract_1d(frames, cy, half_height=3)\n    return spec","metadata":{},"outputs":[],"execution_count":null},{"id":"8a988c7b","cell_type":"code","source":"def build_design_matrix(lightcurve: np.ndarray, centroids: Optional[Tuple[np.ndarray,np.ndarray]]=None, background: Optional[np.ndarray]=None, order: int = 2, pld_pixels: Optional[np.ndarray]=None, pld_max_pixels: int = 50) -> np.ndarray:\n    n = len(lightcurve)\n    cols = [np.ones(n)]\n    t = np.linspace(0, 1, n)\n    for k in range(1, order+1):\n        cols.append(t**k)\n    if centroids is not None:\n        cx, cy = centroids\n        cols.extend([cx, cy, cx*cx, cy*cy, cx*cy])\n    if background is not None:\n        cols.append(background)\n    # Pixel Level Decorrelation (PLD): use brightest pixels as regressors\n    if pld_pixels is not None:\n        # pld_pixels: (n, H, W) or (n, P)\n        P = pld_pixels\n        if P.ndim == 3:\n            P = P.reshape(P.shape[0], -1)\n        # Select top-variance pixels to limit regressors\n        var = np.nanvar(P, axis=0)\n        idx = np.argsort(var)[::-1][:pld_max_pixels]\n        Psel = P[:, idx]\n        # Normalize each frame to sum to 1 (standard PLD)\n        denom = np.nansum(Psel, axis=1, keepdims=True) + 1e-12\n        Ppld = (Psel / denom)\n        cols.append(Ppld)\n    X = np.vstack([c if c.ndim == 1 else c.T for c in cols]).T.astype(np.float64)\n    return X","metadata":{},"outputs":[],"execution_count":null},{"id":"c2fcf085","cell_type":"code","source":"# 18) Build submission.csv Matching Sample Format\ndef build_submission(df_pred: pd.DataFrame) -> pd.DataFrame:\n    cols = ['planet_id','instrument','index','mu','sigma']\n    # If sample_submission exists and has required columns, use its ordering\n    sub = meta.get('sample_submission')\n    if sub is not None and {'planet_id','instrument','index'}.issubset(sub.columns):\n        base = sub[['planet_id','instrument','index']].copy()\n        merged = base.merge(df_pred, on=['planet_id','instrument','index'], how='left')\n        merged['mu'] = merged['mu'].fillna(0.0).astype(np.float64)\n        merged['sigma'] = merged['sigma'].fillna(1e-3).astype(np.float64)\n        return merged[cols]\n    # Otherwise, build from predictions and ensure required columns\n    out = df_pred.copy()\n    for c in cols:\n        if c not in out.columns:\n            out[c] = 0.0 if c in ('mu','sigma') else 0\n    # Sort for stability\n    if set(['planet_id','instrument','index']).issubset(out.columns):\n        out = out.sort_values(['planet_id','instrument','index']).reset_index(drop=True)\n    return out[cols]","metadata":{},"outputs":[],"execution_count":null},{"id":"7993905a","cell_type":"code","source":"# 17) Hyperparameters: defaults + load/save helpers\nHP_PATH = WORK_DIR / 'hparams.json'\nDEFAULT_HPARAMS = {\n    'fgs_aperture_r': 5.0,\n    'poly_order': 2,\n    'sigma_clip': 5.0,\n    'airs_binned_len': 282,\n}\ndef load_hparams(path: Path = HP_PATH, defaults: dict = DEFAULT_HPARAMS) -> dict:\n    try:\n        if path.exists():\n            with open(path, 'r') as f:\n                data = json.load(f)\n            # Merge with defaults to ensure missing keys are filled\n            out = defaults.copy()\n            out.update({k: v for k, v in data.items() if k in defaults or True})\n            return out\n    except Exception as e:\n        print('Warning: failed to load hparams:', e)\n    return defaults.copy()\ndef save_hparams(hp: dict, path: Path = HP_PATH) -> None:\n    try:\n        with open(path, 'w') as f:\n            json.dump(hp, f, indent=2)\n    except Exception as e:\n        print('Warning: failed to save hparams:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"04063d4e","cell_type":"code","source":"# 17b) Bootstrap helpers if running cells out of order\ntry:\n    _ = load_adc_info\nexcept NameError:\n    def load_adc_info(meta: dict) -> tuple[float, float]:\n        df = meta.get('adc_info')\n        if df is None:\n            return 1.0, 0.0\n        gain = float(df.loc[0, 'gain']) if 'gain' in df.columns else 1.0\n        offset = float(df.loc[0, 'offset']) if 'offset' in df.columns else 0.0\n        return gain, offset\n    print('Defined minimal load_adc_info() bootstrap.')","metadata":{},"outputs":[],"execution_count":null},{"id":"2dbf17f3","cell_type":"code","source":"# 21a) Minimal run_inference stub (baseline)\n# This placeholder mirrors sample_submission rows for the selected planets when possible,\n# otherwise it generates a tiny default set of rows so the pipeline can complete end-to-end.\n# Replace with the full extraction/detrending-based inference when ready.\ndef run_inference(split: str, planets: List[str]) -> pd.DataFrame:\n    cols = ['planet_id','instrument','index','mu','sigma']\n    sub = meta.get('sample_submission')\n    # Case A: sample_submission has the expected key columns; mirror it for requested planets\n    if sub is not None and {'planet_id','instrument','index'}.issubset(set(sub.columns)):\n        planet_keys = [str(p) for p in planets]\n        mask = sub['planet_id'].astype(str).isin(planet_keys)\n        base = sub.loc[mask, ['planet_id','instrument','index']].copy()\n        if base.empty:\n            # If none of the requested planets are in sample, fallback to all\n            base = sub[['planet_id','instrument','index']].copy()\n        base['mu'] = 0.0\n        base['sigma'] = 1e-3\n        # Preserve dtypes of key columns where possible\n        try:\n            base['planet_id'] = base['planet_id'].astype(sub['planet_id'].dtype)\n            base['instrument'] = base['instrument'].astype(sub['instrument'].dtype)\n            base['index'] = base['index'].astype(sub['index'].dtype)\n        except Exception:\n            pass\n        return base\n    # Case B: sample_submission exists but lacks required columns OR sub is None\n    # Build a minimal default skeleton: two instruments x one index per requested planet\n    rows = []\n    default_instruments = ['FGS1', 'AIRS-CH0']\n    for pid in planets:\n        for instr in default_instruments:\n            rows.append({\n                'planet_id': pid,\n                'instrument': instr,\n                'index': 0,\n                'mu': 0.0,\n                'sigma': 1e-3,\n            })\n    df = pd.DataFrame(rows, columns=cols)\n    return df","metadata":{},"outputs":[],"execution_count":null},{"id":"806e5df8","cell_type":"code","source":"# 20) Performance and Memory Controls (chunking, caching)\nCONFIG = {\n    'batch_rows_fgs': 5000,\n    'batch_rows_airs': 1024,\n    'use_cache': True\n}\nprint(\"Config:\", CONFIG)\n\n# 21) Reproducibility and Logging + Main pipeline\ndef log_env():\n    # Import locally to avoid issues if global names are shadowed by variables elsewhere\n    import sys as _sys\n    import platform as _platform\n    import numpy as _np\n    import pandas as _pd\n    print(\"Python:\", _sys.version)\n    print(\"Platform:\", _platform.platform())\n    print(\"Numpy:\", _np.__version__, \"Pandas:\", _pd.__version__)\n    try:\n        import pyarrow as _pa\n        print(\"PyArrow:\", _pa.__version__)\n    except Exception:\n        pass\nlog_env()\n\ndef list_planets(split: str) -> List[str]:\n    base = INPUT_DIR / split\n    if not base.exists(): return []\n    return sorted([p.name for p in base.iterdir() if p.is_dir()])\n\n# Main execution: try test set; if unavailable, do a small local dry run\nhp = load_hparams()\nprint(\"Hyperparameters:\", hp)\n\nadc_gain, adc_offset = load_adc_info(meta)\nprint(\"ADC gain/offset:\", adc_gain, adc_offset)\n\ntest_planets = list_planets('test')\nif len(test_planets)==0:\n    print(\"No test planets found; attempting train planets for smoke test...\")\n    test_planets = list_planets('train')[:1]  # limit to 1 for quick local check\n\nif len(test_planets)>0:\n    preds = run_inference('test' if (INPUT_DIR / 'test').exists() else 'train', test_planets)\n    if len(preds)>0:\n        sub = build_submission(preds)\n        sub.to_csv(SUBMISSION_PATH, index=False)\n        print(\"Wrote:\", SUBMISSION_PATH)\n        display(sub.head())\n    else:\n        print(\"No predictions generated.\")\nelse:\n    # fallback sample submission if available\n    if meta['sample_submission'] is not None:\n        meta['sample_submission'][['planet_id','instrument','index']].assign(mu=0.0, sigma=1e-3).to_csv(SUBMISSION_PATH, index=False)\n        print(\"No data found. Wrote empty baseline submission:\", SUBMISSION_PATH)\n    else:\n        print(\"No data or sample submission found — nothing to write.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"d11d9782","cell_type":"code","source":"# Utilities: parquet iteration, FGS helpers, calibration stubs, GLS, time/wavelength\nimport numpy as _np\nimport pandas as _pd\nfrom typing import Iterator as _Iterator, Optional as _Optional, Tuple as _Tuple\n\n# Stream parquet rows in manageable chunks; uses PyArrow when available, else Pandas fallback\ndef iter_parquet_rows(path: Path, batch_rows: int = 2000) -> _Iterator[_np.ndarray]:\n    path = Path(path)\n    if not path.exists():\n        raise FileNotFoundError(str(path))\n    # Prefer PyArrow for large files\n    if pq is not None:\n        try:\n            pf = pq.ParquetFile(str(path))\n            for rg in range(pf.num_row_groups):\n                tbl = pf.read_row_group(rg)\n                df = tbl.to_pandas()\n                arr = df.to_numpy()\n                # If cells are arrays/series per-row, try to expand\n                if arr.dtype == object and hasattr(arr[0, 0], '__len__'):\n                    arr = _np.vstack(arr[:, 0]).astype(_np.float64)\n                for i in range(0, len(arr), batch_rows):\n                    yield _np.ascontiguousarray(arr[i:i+batch_rows], dtype=_np.float64)\n            return\n        except Exception:\n            pass\n    # Pandas fallback (loads whole file then chunks)\n    df = _pd.read_parquet(str(path))\n    arr = df.to_numpy()\n    if arr.dtype == object and hasattr(arr[0, 0], '__len__'):\n        arr = _np.vstack(arr[:, 0]).astype(_np.float64)\n    for i in range(0, len(arr), batch_rows):\n        yield _np.ascontiguousarray(arr[i:i+batch_rows], dtype=_np.float64)\n\n# FGS reshape helper: (n, 1024) -> (n, 32, 32)\ndef fgs_unflatten(arr: _np.ndarray) -> _np.ndarray:\n    arr = _np.asarray(arr)\n    if arr.ndim != 2:\n        raise ValueError('Expected 2D array (n, P)')\n    n, P = arr.shape\n    if P == 1024:\n        return arr.reshape(n, 32, 32)\n    # Heuristic fallback: try square\n    s = int(round(P ** 0.5))\n    if s * s == P:\n        return arr.reshape(n, s, s)\n    # Last resort: assume (H=32, W=P//32)\n    H = 32\n    W = P // H\n    if H * W != P:\n        raise ValueError(f'Cannot reshape FGS array of length {P}')\n    return arr.reshape(n, H, W)\n\n# Simple centroid per frame (weighted by positive flux)\ndef compute_centroids(frames: _np.ndarray) -> _Tuple[_np.ndarray, _np.ndarray]:\n    f3 = _np.asarray(frames, dtype=_np.float64)\n    n, H, W = f3.shape\n    yy, xx = _np.indices((H, W))\n    # Shift to non-negative weights\n    f = f3 - _np.nanmin(f3, axis=(1, 2), keepdims=True)\n    f = _np.clip(f, 0, None)\n    denom = _np.nansum(f, axis=(1, 2)) + 1e-12\n    cx = _np.nansum(f * xx[None, :, :], axis=(1, 2)) / denom\n    cy = _np.nansum(f * yy[None, :, :], axis=(1, 2)) / denom\n    return cx, cy\n\n# Aperture photometry around per-frame centroid\ndef aperture_photometry(frames: _np.ndarray, r: float = 5.0) -> _np.ndarray:\n    f3 = _np.asarray(frames, dtype=_np.float64)\n    n, H, W = f3.shape\n    cx, cy = compute_centroids(f3)\n    yy, xx = _np.indices((H, W))\n    out = _np.empty(n, dtype=_np.float64)\n    rr2 = None\n    for i in range(n):\n        if rr2 is None or True:\n            rr2 = (xx - cx[i])**2 + (yy - cy[i])**2\n        mask = rr2 <= (r*r)\n        out[i] = _np.nansum(f3[i][mask])\n    return out\n\n# Minimal calibration pipeline: ADC restore only (gain/offset). Masters ignored for now.\ndef load_build_calibration(split: str, planet_id: str, instrument: str, gain: float, offset: float) -> dict:\n    return {}\n\ndef calibrate_batch_frames(batch: _np.ndarray, masters: dict, gain: float, offset: float) -> _np.ndarray:\n    arr = _np.asarray(batch, dtype=_np.float64)\n    g = gain if gain not in (None, 0) else 1.0\n    return (arr - (offset or 0.0)) / g\n\n# Simple GLS (OLS) solver\ndef generalized_least_squares(y: _np.ndarray, X: _np.ndarray) -> _Tuple[_np.ndarray, _np.ndarray, float]:\n    y = _np.asarray(y, dtype=_np.float64)\n    X = _np.asarray(X, dtype=_np.float64)\n    # Mask non-finite rows\n    m = _np.isfinite(y) & _np.all(_np.isfinite(X), axis=1)\n    if not _np.any(m):\n        beta = _np.zeros(X.shape[1], dtype=_np.float64)\n        yhat = _np.full_like(y, _np.nan, dtype=_np.float64)\n        return beta, yhat, _np.nan\n    beta, *_ = _np.linalg.lstsq(X[m], y[m], rcond=None)\n    yhat = X @ beta\n    resid = y[m] - (X[m] @ beta)\n    s2 = float(_np.nanvar(resid))\n    return beta, yhat, s2\n\n# Wavelength/time helpers\ndef get_wavelength_grid(meta: dict, instrument: str) -> _Optional[_np.ndarray]:\n    df = meta.get('wavelengths')\n    if df is None:\n        return None\n    try:\n        d = df[df['instrument'].astype(str) == str(instrument)].copy()\n        if 'index' in d.columns:\n            d = d.sort_values('index')\n        if 'wavelength' in d.columns:\n            return d['wavelength'].to_numpy(dtype=_np.float64)\n    except Exception:\n        pass\n    return None\n\ndef build_time_axis(meta: dict, instrument: str, n: int) -> _np.ndarray:\n    # If axis_info contains a suitable cadence, we could integrate it; for robustness, use 0..1.\n    if n <= 1:\n        return _np.zeros(n, dtype=_np.float64)\n    return _np.linspace(0.0, 1.0, int(n))","metadata":{},"outputs":[],"execution_count":null},{"id":"7a454e91","cell_type":"code","source":"# Visualization: AIRS median image and trace profile\ntry:\n    planets = [p for p in (INPUT_DIR / ('test' if (INPUT_DIR / 'test').exists() else 'train')).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        p = signal_path('test' if (INPUT_DIR / 'test').exists() else 'train', pid, 'AIRS-CH0', 0)\n        if p.exists():\n            gain, offset = load_adc_info(meta)\n            masters = load_build_calibration('test' if (INPUT_DIR / 'test').exists() else 'train', pid, 'AIRS-CH0', gain, offset)\n            # take a small batch\n            for batch in iter_parquet_rows(p, batch_rows=256):\n                cal = calibrate_batch_frames(batch, masters, gain, offset)\n                frames = airs_unflatten(cal)\n                med = np.nanmedian(frames, axis=0)\n                prof_y = np.nansum(med, axis=1)\n                prof_x = np.nansum(med, axis=0)\n                fig, axs = plt.subplots(1,3, figsize=(14,4))\n                im = axs[0].imshow(med, aspect='auto', origin='lower')\n                axs[0].set_title(f'AIRS median image (planet {pid})')\n                plt.colorbar(im, ax=axs[0], fraction=0.046, pad=0.04)\n                axs[1].plot(prof_y)\n                axs[1].set_title('Cross-dispersion profile (sum over x)')\n                axs[2].plot(prof_x)\n                axs[2].set_title('Dispersion profile (sum over y)')\n                plt.tight_layout()\n                break\nexcept Exception as e:\n    print('AIRS visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"1e096f5a","cell_type":"code","source":"# Visualization: FGS1 light curve preview\ntry:\n    planets = [p for p in (INPUT_DIR / ('test' if (INPUT_DIR / 'test').exists() else 'train')).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        p = signal_path('test' if (INPUT_DIR / 'test').exists() else 'train', pid, 'FGS1', 0)\n        if p.exists():\n            gain, offset = load_adc_info(meta)\n            masters = load_build_calibration('test' if (INPUT_DIR / 'test').exists() else 'train', pid, 'FGS1', gain, offset)\n            vals = []\n            for batch in iter_parquet_rows(p, batch_rows=5000):\n                cal = calibrate_batch_frames(batch, masters, gain, offset)\n                f3 = fgs_unflatten(cal)\n                lc = aperture_photometry(f3, r=5.0)\n                vals.append(lc)\n                if len(np.concatenate(vals)) > 5000:\n                    break\n            if vals:\n                lc_full = np.concatenate(vals)\n                t = build_time_axis(meta, 'FGS1', len(lc_full))\n                plt.figure(figsize=(12,3))\n                plt.plot(t, lc_full/np.nanmedian(lc_full), '-', lw=0.5)\n                plt.title(f'FGS1 light curve (planet {pid})')\n                plt.xlabel('time')\n                plt.ylabel('normalized flux')\n                plt.tight_layout()\nexcept Exception as e:\n    print('FGS1 visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"71906285","cell_type":"code","source":"# Visualization: Quick AIRS transmission spectrum (depth vs wavelength)\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    planets = [p for p in (INPUT_DIR / split).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        p = signal_path(split, pid, 'AIRS-CH0', 0)\n        if p.exists():\n            # Load helpers\n            gain, offset = load_adc_info(meta)\n            masters = load_build_calibration(split, pid, 'AIRS-CH0', gain, offset)\n            # Stream a limited number of frames for speed\n            frames = []\n            max_frames = 1200\n            for batch in iter_parquet_rows(p, batch_rows=2000):\n                cal = calibrate_batch_frames(batch, masters, gain, offset)\n                f3 = airs_unflatten(cal)  # (n, 32, 356)\n                # Apply same crop used in extraction for cleaner spectra\n                f3 = crop_airs_x(f3, 39, 321)\n                frames.append(f3)\n                if sum(x.shape[0] for x in frames) >= max_frames:\n                    break\n            if frames:\n                cube = np.concatenate(frames, axis=0)  # (T, 32, W)\n                T = cube.shape[0]\n                # Collapse cross-dispersion to 1D spectra (simple sum; we already calibrated)\n                spec = cube.sum(axis=1)  # (T, W)\n                # Build time axis and choose a naive transit window (middle 20% as in, edges 40% as out)\n                t = build_time_axis(meta, 'AIRS-CH0', T)\n                i0, i1 = int(0.4*T), int(0.6*T)\n                in_m = np.zeros(T, dtype=bool)\n                in_m[i0:i1] = True\n                out_m = ~in_m\n                # Normalize each wavelength by its out-of-transit median\n                out_med = np.nanmedian(spec[out_m], axis=0)\n                out_med[out_med==0] = np.nan\n                norm = spec / out_med[np.newaxis, :]\n                # Depth per wavelength: 1 - mean_in\n                mean_in = np.nanmean(norm[in_m], axis=0)\n                mean_out = np.nanmean(norm[out_m], axis=0)\n                var_in = np.nanvar(norm[in_m], axis=0)/(in_m.sum() + 1e-9)\n                var_out = np.nanvar(norm[out_m], axis=0)/(out_m.sum() + 1e-9)\n                depth = 1.0 - (mean_in/mean_out)\n                sigma = np.sqrt(var_in + var_out)\n                # Optional wavelength grid (crop-aware if full grid available)\n                waves = None\n                try:\n                    w = get_wavelength_grid(meta, 'AIRS-CH0')\n                    if w is not None:\n                        w = w[39:321]\n                        if len(w) == spec.shape[1]:\n                            waves = w\n                except Exception:\n                    pass\n                plt.figure(figsize=(12,4))\n                x = np.arange(spec.shape[1]) if waves is None else waves\n                plt.plot(x, depth, color='tab:blue', lw=1)\n                if np.isfinite(sigma).any():\n                    lo = depth - sigma\n                    hi = depth + sigma\n                    plt.fill_between(x, lo, hi, color='tab:blue', alpha=0.2, linewidth=0)\n                plt.title(f'AIRS quick spectrum (planet {pid}, visit 0)')\n                plt.xlabel('wavelength' if waves is not None else 'pixel column (cropped)')\n                plt.ylabel('transit depth (arb)')\n                plt.tight_layout()\n                # Save figure\n                fig_dir = WORK_DIR / 'figures'\n                fig_dir.mkdir(parents=True, exist_ok=True)\n                out_path = fig_dir / f'{pid}_AIRS-CH0_quick_spectrum.png'\n                plt.savefig(out_path, dpi=150)\n                print('Saved:', out_path)\nexcept Exception as e:\n    print('AIRS spectrum visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"703196c6","cell_type":"code","source":"# Visualization: FGS centroids + raw vs detrended, and save figures\nfrom pathlib import Path\nfig_dir = WORK_DIR / 'figures'\nfig_dir.mkdir(parents=True, exist_ok=True)\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    planets = [p for p in (INPUT_DIR / split).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        p = signal_path(split, pid, 'FGS1', 0)\n        if p.exists():\n            # Load calibration\n            gain, offset = load_adc_info(meta)\n            masters = load_build_calibration(split, pid, 'FGS1', gain, offset)\n            # Stream frames and compute time, flux, centroids, and keep small cube for PLD\n            flux_chunks, cx_chunks, cy_chunks = [], [], []\n            cubes = []\n            cap_frames = 4000\n            for batch in iter_parquet_rows(p, batch_rows=4000):\n                cal = calibrate_batch_frames(batch, masters, gain, offset)\n                f3 = fgs_unflatten(cal)\n                flux_chunks.append(aperture_photometry(f3, r=5.0))\n                cx, cy = compute_centroids(f3)\n                cx_chunks.append(cx)\n                cy_chunks.append(cy)\n                if sum(len(a) for a in flux_chunks) < cap_frames:\n                    cubes.append(f3)\n                if sum(len(a) for a in flux_chunks) >= 8000:\n                    break\n            if flux_chunks:\n                flux = np.concatenate(flux_chunks)\n                cx = np.concatenate(cx_chunks)\n                cy = np.concatenate(cy_chunks)\n                t = build_time_axis(meta, 'FGS1', len(flux))\n                # Baseline detrend with centroids + poly ramps\n                X = build_design_matrix(flux, centroids=(cx, cy), order=2)\n                beta, yhat, s2 = generalized_least_squares(flux, X)\n                flux_detr = flux - yhat + np.nanmedian(flux)\n                # Optional PLD detrend using early frames cube\n                flux_pld = None\n                if cubes:\n                    cube_small = np.concatenate(cubes, axis=0)\n                    n_use = min(len(flux), cube_small.shape[0])\n                    X_pld = build_design_matrix(flux[:n_use], centroids=(cx[:n_use], cy[:n_use]), order=2, pld_pixels=cube_small[:n_use], pld_max_pixels=50)\n                    _, yhat_pld, _ = generalized_least_squares(flux[:n_use], X_pld)\n                    flux_pld = flux[:n_use] - yhat_pld + np.nanmedian(flux[:n_use])\n                # Plot centroids vs time\n                plt.figure(figsize=(12,3))\n                plt.plot(t, cx, label='centroid_x', lw=0.6)\n                plt.plot(t, cy, label='centroid_y', lw=0.6)\n                plt.legend()\n                plt.title(f'FGS1 centroids (planet {pid})')\n                plt.xlabel('time'); plt.ylabel('pixels')\n                plt.tight_layout()\n                out_path1 = fig_dir / f'{pid}_FGS1_centroids.png'\n                plt.savefig(out_path1, dpi=150)\n                # Plot raw vs detrended flux\n                plt.figure(figsize=(12,3))\n                nf = flux/np.nanmedian(flux)\n                nd = flux_detr/np.nanmedian(flux_detr)\n                plt.plot(t, nf, label='raw', alpha=0.7, lw=0.6)\n                plt.plot(t, nd, label='detrended', alpha=0.8, lw=0.8)\n                if flux_pld is not None:\n                    tt = t[:len(flux_pld)]\n                    npd = flux_pld/np.nanmedian(flux_pld)\n                    plt.plot(tt, npd, label='detrended+PLD', alpha=0.9, lw=0.8)\n                plt.legend()\n                plt.title(f'FGS1 raw vs detrended (planet {pid})')\n                plt.xlabel('time'); plt.ylabel('normalized flux')\n                plt.tight_layout()\n                out_path2 = fig_dir / f'{pid}_FGS1_raw_vs_detrended.png'\n                plt.savefig(out_path2, dpi=150)\n                print('Saved:', out_path1)\n                print('Saved:', out_path2)\nexcept Exception as e:\n    print('FGS centroids/detrend visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"04b176d3","cell_type":"code","source":"# Visualization: AIRS spectrogram (time vs wavelength) and save\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    planets = [p for p in (INPUT_DIR / split).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        p = signal_path(split, pid, 'AIRS-CH0', 0)\n        if p.exists():\n            gain, offset = load_adc_info(meta)\n            masters = load_build_calibration(split, pid, 'AIRS-CH0', gain, offset)\n            frames = []\n            max_frames = 1500\n            for batch in iter_parquet_rows(p, batch_rows=1500):\n                cal = calibrate_batch_frames(batch, masters, gain, offset)\n                f3 = airs_unflatten(cal)\n                f3 = crop_airs_x(f3, 39, 321)\n                frames.append(f3)\n                if sum(x.shape[0] for x in frames) >= max_frames:\n                    break\n            if frames:\n                cube = np.concatenate(frames, axis=0)  # (T, 32, W)\n                T, H, W = cube.shape\n                spec = cube.sum(axis=1)  # (T, W)\n                t = build_time_axis(meta, 'AIRS-CH0', T)\n                # Naive in/out masks (middle 20% in-transit)\n                i0, i1 = int(0.4*T), int(0.6*T)\n                in_m = np.zeros(T, dtype=bool)\n                in_m[i0:i1] = True\n                out_m = ~in_m\n                out_med = np.nanmedian(spec[out_m], axis=0)\n                out_med[out_med==0] = np.nan\n                norm = spec / out_med[np.newaxis, :]\n                spect = norm - 1.0\n                # Wavelength grid (cropped)\n                waves = None\n                try:\n                    w = get_wavelength_grid(meta, 'AIRS-CH0')\n                    if w is not None:\n                        w = w[39:321]\n                        if len(w) == W:\n                            waves = w\n                except Exception:\n                    pass\n                # Plot and save\n                fig, ax = plt.subplots(1,1, figsize=(12,4))\n                extent = [0, W-1, t[0], t[-1]] if waves is None else [waves[0], waves[-1], t[0], t[-1]]\n                im = ax.imshow(spect, aspect='auto', origin='lower', extent=extent, cmap='coolwarm', vmin=-0.02, vmax=0.02)\n                ax.set_xlabel('wavelength' if waves is not None else 'pixel column (cropped)')\n                ax.set_ylabel('time')\n                ax.set_title(f'AIRS spectrogram (planet {pid})')\n                plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04, label='(flux/out) - 1')\n                fig_dir = WORK_DIR / 'figures'\n                fig_dir.mkdir(parents=True, exist_ok=True)\n                out_path = fig_dir / f'{pid}_AIRS-CH0_spectrogram.png'\n                plt.tight_layout()\n                plt.savefig(out_path, dpi=150)\n                print('Saved:', out_path)\nexcept Exception as e:\n    print('AIRS spectrogram visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"92bdbb5a","cell_type":"code","source":"# Visualization: FGS diagnostics — centroid scatter and PSD; save\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    planets = [p for p in (INPUT_DIR / split).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        p = signal_path(split, pid, 'FGS1', 0)\n        if p.exists():\n            gain, offset = load_adc_info(meta)\n            masters = load_build_calibration(split, pid, 'FGS1', gain, offset)\n            flux_chunks, cx_chunks, cy_chunks = [], [], []\n            for batch in iter_parquet_rows(p, batch_rows=4000):\n                cal = calibrate_batch_frames(batch, masters, gain, offset)\n                f3 = fgs_unflatten(cal)\n                flux_chunks.append(aperture_photometry(f3, r=5.0))\n                cx, cy = compute_centroids(f3)\n                cx_chunks.append(cx)\n                cy_chunks.append(cy)\n                if sum(len(a) for a in flux_chunks) >= 8000:\n                    break\n            if flux_chunks:\n                flux = np.concatenate(flux_chunks)\n                cx = np.concatenate(cx_chunks)\n                cy = np.concatenate(cy_chunks)\n                t = build_time_axis(meta, 'FGS1', len(flux))\n                X = build_design_matrix(flux, centroids=(cx, cy), order=2)\n                _, yhat, _ = generalized_least_squares(flux, X)\n                detr = flux - yhat + np.nanmedian(flux)\n                # Centroid scatter\n                plt.figure(figsize=(4,4))\n                plt.scatter(cx, cy, s=1, alpha=0.5)\n                plt.xlabel('centroid_x')\n                plt.ylabel('centroid_y')\n                plt.title(f'FGS centroid scatter (planet {pid})')\n                fig_dir = WORK_DIR / 'figures'\n                fig_dir.mkdir(parents=True, exist_ok=True)\n                out1 = fig_dir / f'{pid}_FGS1_centroid_scatter.png'\n                plt.tight_layout(); plt.savefig(out1, dpi=150)\n                # PSD of raw vs detrended\n                try:\n                    from matplotlib import mlab\n                    fs = 1.0/np.median(np.diff(t)) if len(t)>1 else 1.0\n                    Pxx_raw, f_raw = mlab.psd((flux/np.nanmedian(flux))-1, NFFT=256, Fs=fs)\n                    Pxx_det, f_det = mlab.psd((detr/np.nanmedian(detr))-1, NFFT=256, Fs=fs)\n                    plt.figure(figsize=(6,3))\n                    plt.semilogy(f_raw, Pxx_raw, label='raw')\n                    plt.semilogy(f_det, Pxx_det, label='detrended')\n                    plt.xlabel('frequency (1/time)'); plt.ylabel('PSD')\n                    plt.title(f'FGS PSD (planet {pid})')\n                    plt.legend(); plt.tight_layout()\n                    out2 = fig_dir / f'{pid}_FGS1_psd.png'\n                    plt.savefig(out2, dpi=150)\n                    print('Saved:', out1)\n                    print('Saved:', out2)\n                except Exception:\n                    pass\nexcept Exception as e:\n    print('FGS diagnostics visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"d1fc1cd7","cell_type":"code","source":"# Visualization: AIRS single-channel light curve with naive detrending; save\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    planets = [p for p in (INPUT_DIR / split).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        p = signal_path(split, pid, 'AIRS-CH0', 0)\n        if p.exists():\n            gain, offset = load_adc_info(meta)\n            masters = load_build_calibration(split, pid, 'AIRS-CH0', gain, offset)\n            frames = []\n            for batch in iter_parquet_rows(p, batch_rows=1500):\n                cal = calibrate_batch_frames(batch, masters, gain, offset)\n                f3 = airs_unflatten(cal)\n                f3 = crop_airs_x(f3, 39, 321)\n                frames.append(f3)\n                if sum(x.shape[0] for x in frames) >= 2000:\n                    break\n            if frames:\n                cube = np.concatenate(frames, axis=0)  # (T, 32, W)\n                spec = cube.sum(axis=1)  # (T, W)\n                T, W = spec.shape\n                j = W//2  # mid channel\n                y = spec[:, j]\n                t = build_time_axis(meta, 'AIRS-CH0', T)\n                # Detrend with polynomial time only\n                X = build_design_matrix(y, order=2)\n                _, yhat, _ = generalized_least_squares(y, X)\n                yd = y - yhat + np.nanmedian(y)\n                plt.figure(figsize=(12,3))\n                plt.plot(t, y/np.nanmedian(y), label='raw', lw=0.6)\n                plt.plot(t, yd/np.nanmedian(yd), label='detrended', lw=0.8)\n                plt.legend()\n                plt.xlabel('time'); plt.ylabel('normalized flux')\n                plt.title(f'AIRS single-channel LC (planet {pid}, col {j})')\n                plt.tight_layout()\n                fig_dir = WORK_DIR / 'figures'\n                fig_dir.mkdir(parents=True, exist_ok=True)\n                out = fig_dir / f'{pid}_AIRS-CH0_single_channel_lc.png'\n                plt.savefig(out, dpi=150)\n                print('Saved:', out)\nexcept Exception as e:\n    print('AIRS single-channel LC visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"8d7c8b86","cell_type":"code","source":"# 17c) Path helpers (bootstrap) if running cells out of order\ntry:\n    _ = signal_path  # type: ignore\nexcept NameError:\n    try:\n        SIG_EXT_ = SIG_EXT  # use existing if defined\n    except NameError:\n        SIG_EXT_ = '.parquet'\n    def planet_dir(split: str, planet_id: str) -> Path:\n        return INPUT_DIR / split / str(planet_id)\n    def signal_path(split: str, planet_id: str, instrument: str, visit_idx: int = 0) -> Path:\n        base = planet_dir(split, planet_id)\n        # Try a few common layouts; return first existing, else a sensible default\n        candidates = [\n            base / instrument / f\"visit_{visit_idx:02d}{SIG_EXT_}\",\n            base / instrument / f\"{visit_idx:02d}{SIG_EXT_}\",\n            base / instrument / f\"visit_{visit_idx}{SIG_EXT_}\",\n            base / instrument / f\"{visit_idx}{SIG_EXT_}\",\n            base / f\"{instrument}_visit_{visit_idx:02d}{SIG_EXT_}\",\n            base / f\"{instrument}_{visit_idx:02d}{SIG_EXT_}\",\n        ]\n        for p in candidates:\n            try:\n                if p.exists():\n                    return p\n            except Exception:\n                pass\n        # Fallback: first parquet under instrument folder if any\n        try:\n            inst_dir = base / instrument\n            if inst_dir.exists():\n                for p in inst_dir.glob(f\"*{SIG_EXT_}\"):\n                    return p\n        except Exception:\n            pass\n        return candidates[0]\n    print('Defined minimal signal_path() and planet_dir() bootstrap.')","metadata":{},"outputs":[],"execution_count":null},{"id":"7a51e6b9","cell_type":"code","source":"# Visualization: Dataset overview — planets, instruments, and file counts\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    base = INPUT_DIR / split\n    if not base.exists():\n        raise FileNotFoundError(f\"split folder not found: {base}\")\n    planets = [p for p in base.iterdir() if p.is_dir()]\n    if not planets:\n        raise RuntimeError('No planet directories found')\n    # Collect counts per instrument for first N planets\n    N = 8\n    instruments = ['FGS1','FGS2','AIRS-CH0','AIRS-CH1']\n    rows = []\n    for pid_path in planets[:N]:\n        pid = pid_path.name\n        for inst in instruments:\n            inst_dir = pid_path / inst\n            cnt = 0\n            if inst_dir.exists():\n                try:\n                    cnt = sum(1 for _ in inst_dir.glob('*.parquet'))\n                except Exception:\n                    cnt = 0\n            rows.append({'planet_id': pid, 'instrument': inst, 'files': cnt})\n    import pandas as _pd\n    dfc = _pd.DataFrame(rows)\n    if dfc['files'].sum() == 0:\n        raise RuntimeError('No parquet files found for the first few planets')\n    # Pivot to planet x instrument matrix\n    piv = dfc.pivot(index='planet_id', columns='instrument', values='files').fillna(0)\n    # Plot\n    ax = piv.plot(kind='bar', figsize=(12,4))\n    ax.set_title(f'Dataset overview: file counts per instrument (first {len(piv)} planets in {split})')\n    ax.set_xlabel('planet_id'); ax.set_ylabel('# files')\n    plt.tight_layout()\n    fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n    out = fig_dir / f'dataset_overview_{split}.png'\n    plt.savefig(out, dpi=150)\n    print('Saved:', out)\nexcept Exception as e:\n    print('Dataset overview visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"29071726","cell_type":"code","source":"# Visualization: AIRS trace center drift and width over time\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    planets = [p for p in (INPUT_DIR / split).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        pth = signal_path(split, pid, 'AIRS-CH0', 0)\n        if not pth.exists():\n            raise FileNotFoundError(pth)\n        gain, offset = load_adc_info(meta)\n        masters = load_build_calibration(split, pid, 'AIRS-CH0', gain, offset)\n        cys, widths = [], []\n        Tcap = 2000\n        seen = 0\n        for batch in iter_parquet_rows(pth, batch_rows=512):\n            cal = calibrate_batch_frames(batch, masters, gain, offset)\n            f3 = airs_unflatten(cal)\n            f3 = crop_airs_x(f3, 39, 321)\n            med = np.nanmedian(f3, axis=2)  # collapse x -> shape (n, 32)\n            # center and width per frame using simple weighted stats\n            y = np.arange(32)\n            w = med\n            denom = np.nansum(w, axis=1) + 1e-12\n            cy = np.nansum(w * y[None, :], axis=1) / denom\n            var = np.nansum(w * (y[None, :] - cy[:, None])**2, axis=1) / denom\n            cys.append(cy)\n            widths.append(np.sqrt(np.maximum(var, 0)))\n            seen += len(cy)\n            if seen >= Tcap:\n                break\n        if cys:\n            cy_all = np.concatenate(cys)\n            wd_all = np.concatenate(widths)\n            t = build_time_axis(meta, 'AIRS-CH0', len(cy_all))\n            fig, ax = plt.subplots(2,1, figsize=(12,5), sharex=True)\n            ax[0].plot(t, cy_all, lw=0.6)\n            ax[0].set_ylabel('center_y [px]')\n            ax[0].set_title(f'AIRS trace center drift (planet {pid})')\n            ax[1].plot(t, wd_all, lw=0.6)\n            ax[1].set_ylabel('width [px]'); ax[1].set_xlabel('time')\n            plt.tight_layout()\n            fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n            out = fig_dir / f'{pid}_AIRS-CH0_trace_center_width.png'\n            plt.savefig(out, dpi=150)\n            print('Saved:', out)\nexcept Exception as e:\n    print('AIRS center/width visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"9e17873f","cell_type":"code","source":"# Visualization: FGS aperture growth curve (flux vs aperture radius)\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    planets = [p for p in (INPUT_DIR / split).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        pth = signal_path(split, pid, 'FGS1', 0)\n        if not pth.exists():\n            raise FileNotFoundError(pth)\n        gain, offset = load_adc_info(meta)\n        masters = load_build_calibration(split, pid, 'FGS1', gain, offset)\n        # take one batch of frames\n        batch0 = next(iter_parquet_rows(pth, batch_rows=512))\n        cal = calibrate_batch_frames(batch0, masters, gain, offset)\n        f3 = fgs_unflatten(cal)\n        # median image to define center\n        med = np.nanmedian(f3, axis=0)\n        yy, xx = np.indices(med.shape)\n        # rough centroid on median\n        w = med - np.nanmin(med)\n        w = np.clip(w, 0, None)\n        denom = np.nansum(w) + 1e-12\n        cx = np.nansum(w * xx) / denom\n        cy = np.nansum(w * yy) / denom\n        rs = np.linspace(2, 12, 11)\n        fluxes = []\n        for r in rs:\n            rr = ((xx - cx)**2 + (yy - cy)**2)**0.5\n            mask = rr <= r\n            # integrate over mask for each frame then median across time\n            vals = np.nansum(f3[:, mask], axis=1)\n            fluxes.append(np.nanmedian(vals))\n        plt.figure(figsize=(6,4))\n        plt.plot(rs, fluxes, marker='o')\n        plt.xlabel('aperture radius [px]'); plt.ylabel('median aperture flux')\n        plt.title(f'FGS1 aperture growth curve (planet {pid})')\n        plt.tight_layout()\n        fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n        out = fig_dir / f'{pid}_FGS1_aperture_growth.png'\n        plt.savefig(out, dpi=150)\n        print('Saved:', out)\nexcept Exception as e:\n    print('FGS aperture growth visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"35f423a5","cell_type":"code","source":"# Visualization: FGS design-matrix correlations and residual histogram\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    planets = [p for p in (INPUT_DIR / split).iterdir() if p.is_dir()]\n    if planets:\n        pid = planets[0].name\n        pth = signal_path(split, pid, 'FGS1', 0)\n        if not pth.exists():\n            raise FileNotFoundError(pth)\n        gain, offset = load_adc_info(meta)\n        masters = load_build_calibration(split, pid, 'FGS1', gain, offset)\n        flux_chunks, cx_chunks, cy_chunks, cubes = [], [], [], []\n        cap_frames = 3000\n        for batch in iter_parquet_rows(pth, batch_rows=1500):\n            cal = calibrate_batch_frames(batch, masters, gain, offset)\n            f3 = fgs_unflatten(cal)\n            flux_chunks.append(aperture_photometry(f3, r=5.0))\n            cx, cy = compute_centroids(f3)\n            cx_chunks.append(cx); cy_chunks.append(cy)\n            if sum(len(a) for a in flux_chunks) < cap_frames:\n                cubes.append(f3)\n            if sum(len(a) for a in flux_chunks) >= cap_frames:\n                break\n        if flux_chunks:\n            import numpy as _np\n            flux = _np.concatenate(flux_chunks)\n            cx = _np.concatenate(cx_chunks); cy = _np.concatenate(cy_chunks)\n            t = build_time_axis(meta, 'FGS1', len(flux))\n            # Build design matrix with centroids and PLD from a small cube\n            P = None\n            if cubes:\n                cube_small = _np.concatenate(cubes, axis=0)\n                n_use = min(len(flux), cube_small.shape[0], 2500)\n                P = cube_small[:n_use]\n            n_use = len(flux) if P is None else min(len(flux), P.shape[0])\n            X = build_design_matrix(flux[:n_use], centroids=(cx[:n_use], cy[:n_use]), order=2, pld_pixels=P[:n_use] if P is not None else None, pld_max_pixels=40)\n            # Correlation heatmap\n            C = _np.corrcoef(X.T)\n            plt.figure(figsize=(6,5))\n            im = plt.imshow(C, vmin=-1, vmax=1, cmap='coolwarm')\n            plt.colorbar(im, fraction=0.046, pad=0.04)\n            plt.title(f'FGS design-matrix correlation (planet {pid})')\n            plt.tight_layout()\n            fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n            out1 = fig_dir / f'{pid}_FGS1_design_corr.png'\n            plt.savefig(out1, dpi=150)\n            # Residual histogram\n            beta, yhat, s2 = generalized_least_squares(flux[:n_use], X)\n            resid = flux[:n_use] - yhat\n            plt.figure(figsize=(6,3))\n            plt.hist((resid - _np.nanmedian(resid))/_np.nanstd(resid), bins=60, alpha=0.8)\n            plt.title(f'FGS residuals (z-scored), planet {pid}')\n            plt.xlabel('z'); plt.ylabel('count'); plt.tight_layout()\n            out2 = fig_dir / f'{pid}_FGS1_residual_hist.png'\n            plt.savefig(out2, dpi=150)\n            print('Saved:', out1)\n            print('Saved:', out2)\nexcept Exception as e:\n    print('FGS design-matrix/residual visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"80c400e9","cell_type":"code","source":"# Visualization: Global plotting style for publication-quality figures\ntry:\n    import matplotlib as mpl\n    mpl.rcParams.update({\n        'figure.dpi': 120,\n        'savefig.dpi': 150,\n        'axes.grid': True,\n        'grid.alpha': 0.25,\n        'axes.spines.top': False,\n        'axes.spines.right': False,\n        'axes.titlesize': 12,\n        'axes.labelsize': 11,\n        'legend.fontsize': 9,\n        'xtick.labelsize': 9,\n        'ytick.labelsize': 9,\n        'image.cmap': 'cividis',\n    })\n    if sns is not None:\n        sns.set_theme(style='whitegrid', context='notebook')\n    print('Plotting style configured.')\nexcept Exception as e:\n    print('Style configuration skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"071c325c","cell_type":"code","source":"# Advanced: FGS diagnostics dashboard (raw/detrended, centroids, residuals, RMS)\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    base = INPUT_DIR / split\n    planets = [p for p in base.iterdir() if p.is_dir()] if base.exists() else []\n    pid = planets[0].name if planets else 'SYNTH'\n    have_data = bool(planets)\n    if have_data:\n        pth = signal_path(split, pid, 'FGS1', 0)\n        gain, offset = load_adc_info(meta)\n        masters = load_build_calibration(split, pid, 'FGS1', gain, offset)\n        flux_chunks, cx_chunks, cy_chunks, cubes = [], [], [], []\n        for batch in iter_parquet_rows(pth, batch_rows=4000):\n            cal = calibrate_batch_frames(batch, masters, gain, offset)\n            f3 = fgs_unflatten(cal)\n            flux_chunks.append(aperture_photometry(f3, r=HP_PATH.exists() and load_hparams().get('fgs_aperture_r',5.0) or 5.0))\n            cx, cy = compute_centroids(f3)\n            cx_chunks.append(cx); cy_chunks.append(cy)\n            cubes.append(f3)\n            if sum(len(a) for a in flux_chunks) >= 12000:\n                break\n        flux = np.concatenate(flux_chunks) if flux_chunks else None\n        cx = np.concatenate(cx_chunks) if cx_chunks else None\n        cy = np.concatenate(cy_chunks) if cy_chunks else None\n        cube_small = np.concatenate(cubes, axis=0) if cubes else None\n    else:\n        # Synthetic fallback\n        n = 3000\n        t = np.linspace(0, 1, n)\n        cx = 16 + 0.1*np.sin(2*np.pi*3*t) + 0.02*np.random.randn(n)\n        cy = 16 + 0.1*np.cos(2*np.pi*2*t) + 0.02*np.random.randn(n)\n        sys = 1 + 0.005*t + 0.002*np.sin(2*np.pi*5*t) + 0.001*cx + 0.001*cy\n        transit = 1 - 0.005*(np.abs(t-0.5)<0.05)\n        flux = sys * transit * (1 + 0.001*np.random.randn(n))\n        cube_small = None\n    if flux is None or len(flux) < 50:\n        raise RuntimeError('Insufficient FGS data for dashboard')\n    t = build_time_axis(meta, 'FGS1', len(flux)) if have_data else t\n    # Build design matrix and detrend\n    X = build_design_matrix(flux, centroids=(cx, cy), order=2, pld_pixels=cube_small[:len(flux)] if (cube_small is not None and cube_small.shape[0] >= len(flux)) else None, pld_max_pixels=60)\n    beta, yhat, s2 = generalized_least_squares(flux, X)\n    nf = flux/np.nanmedian(flux)\n    nd = (flux - yhat + np.nanmedian(flux))/np.nanmedian(flux)\n    # Time-averaging RMS (Allan-like)\n    def time_avg_rms(y, max_bin=100):\n        y = y - np.nanmedian(y)\n        rms_x, rms_y = [], []\n        for b in np.unique(np.logspace(0, np.log10(max_bin), 20).astype(int)):\n            if b < 1: b = 1\n            m = len(y)//b\n            if m < 2: break\n            yy = y[:m*b].reshape(m, b).mean(axis=1)\n            rms_x.append(b)\n            rms_y.append(np.nanstd(yy))\n        return np.array(rms_x), np.array(rms_y)\n    bx_raw, by_raw = time_avg_rms(nf-1)\n    bx_det, by_det = time_avg_rms(nd-1)\n    # Figure layout\n    fig = plt.figure(figsize=(14,9))\n    gs = fig.add_gridspec(3, 2, height_ratios=[2,1,1], hspace=0.35, wspace=0.25)\n    ax1 = fig.add_subplot(gs[0, :])\n    ax2 = fig.add_subplot(gs[1, 0])\n    ax3 = fig.add_subplot(gs[1, 1])\n    ax4 = fig.add_subplot(gs[2, 0])\n    ax5 = fig.add_subplot(gs[2, 1])\n    # Panel 1: raw vs detrended\n    ax1.plot(t, nf, lw=0.5, alpha=0.8, label='raw')\n    ax1.plot(t, nd, lw=0.8, alpha=0.9, label='detrended')\n    ax1.set_title(f'FGS1 flux: raw vs detrended ({pid})')\n    ax1.set_xlabel('time'); ax1.set_ylabel('normalized flux'); ax1.legend(loc='best')\n    # Panel 2/3: centroids\n    ax2.plot(t, cx, lw=0.6, color='tab:green')\n    ax2.set_title('Centroid X'); ax2.set_xlabel('time'); ax2.set_ylabel('px')\n    ax3.plot(t, cy, lw=0.6, color='tab:orange')\n    ax3.set_title('Centroid Y'); ax3.set_xlabel('time'); ax3.set_ylabel('px')\n    # Panel 4: residual histogram\n    resid = (nf - nd)\n    ax4.hist((resid - np.nanmedian(resid))/ (np.nanstd(resid)+1e-12), bins=60, color='tab:blue', alpha=0.8)\n    ax4.set_title('Residuals (z-scored)'); ax4.set_xlabel('z'); ax4.set_ylabel('count')\n    # Panel 5: time-averaging RMS\n    ax5.loglog(bx_raw, by_raw, 'o-', label='raw')\n    ax5.loglog(bx_det, by_det, 'o-', label='detrended')\n    ax5.set_title('Time-averaging RMS'); ax5.set_xlabel('bin size'); ax5.set_ylabel('RMS')\n    ax5.legend()\n    plt.tight_layout()\n    fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n    out = fig_dir / f'{pid}_FGS1_dashboard.png'\n    plt.savefig(out)\n    print('Saved:', out)\nexcept Exception as e:\n    print('FGS diagnostics dashboard skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"448f9353","cell_type":"code","source":"# Advanced: AIRS spectral dashboard (spectrogram, depth, rolling bands)\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    base = INPUT_DIR / split\n    planets = [p for p in base.iterdir() if p.is_dir()] if base.exists() else []\n    pid = planets[0].name if planets else 'SYNTH'\n    have_data = bool(planets)\n    if have_data:\n        pth = signal_path(split, pid, 'AIRS-CH0', 0)\n        gain, offset = load_adc_info(meta)\n        masters = load_build_calibration(split, pid, 'AIRS-CH0', gain, offset)\n        frames = []\n        cap = 2000\n        for batch in iter_parquet_rows(pth, batch_rows=1000):\n            cal = calibrate_batch_frames(batch, masters, gain, offset)\n            f3 = airs_unflatten(cal)\n            f3 = crop_airs_x(f3, 39, 321)\n            frames.append(f3)\n            if sum(x.shape[0] for x in frames) >= cap:\n                break\n        if frames:\n            cube = np.concatenate(frames, axis=0)\n        else:\n            cube = None\n    else:\n        # Synthetic AIRS cube: transit-like dip in middle across modest wavelength band\n        T, H, W = 800, 32, 260\n        t = np.linspace(0, 1, T)\n        waves = np.linspace(1.1, 1.9, W)\n        band = np.exp(-0.5*((waves-1.5)/0.15)**2)\n        sys = 1 + 0.01*np.sin(2*np.pi*3*t)[:,None,None]\n        transit = 1 - (0.01*band[None,None,:]) * (np.abs(t-0.5)<0.06)[:,None,None]\n        cube = sys*transit*(1 + 0.001*np.random.randn(T,H,W))\n    if cube is None or cube.shape[0] < 50:\n        raise RuntimeError('Insufficient AIRS data for dashboard')\n    T, H, W = cube.shape\n    spec = cube.sum(axis=1)  # (T, W)\n    t = build_time_axis(meta, 'AIRS-CH0', T) if have_data else t\n    # In/Out masks: middle 20% as in-transit\n    i0, i1 = int(0.4*T), int(0.6*T)\n    in_m = np.zeros(T, dtype=bool); in_m[i0:i1] = True\n    out_m = ~in_m\n    out_med = np.nanmedian(spec[out_m], axis=0)\n    out_med[out_med==0] = np.nan\n    norm = spec / out_med[np.newaxis, :]\n    # Depth and uncertainty\n    mean_in = np.nanmean(norm[in_m], axis=0)\n    mean_out = np.nanmean(norm[out_m], axis=0)\n    var_in = np.nanvar(norm[in_m], axis=0)/(in_m.sum() + 1e-9)\n    var_out = np.nanvar(norm[out_m], axis=0)/(out_m.sum() + 1e-9)\n    depth = 1.0 - (mean_in/mean_out)\n    sigma = np.sqrt(np.maximum(var_in + var_out, 0))\n    # Rolling bands time series: split W into 6 equal bands\n    B = 6\n    edges = np.linspace(0, W, B+1).astype(int)\n    band_series = []\n    for i in range(B):\n        a,b = edges[i], edges[i+1]\n        y = np.nanmean(norm[:, a:b], axis=1)\n        band_series.append((a,b,y))\n    # Figure layout\n    fig = plt.figure(figsize=(14,9))\n    gs = fig.add_gridspec(3, 2, height_ratios=[2,1,1], hspace=0.35, wspace=0.25)\n    ax1 = fig.add_subplot(gs[0, :])\n    ax2 = fig.add_subplot(gs[1, 0])\n    ax3 = fig.add_subplot(gs[1, 1])\n    ax4 = fig.add_subplot(gs[2, :])\n    # Panel 1: Spectrogram (normalized)\n    im = ax1.imshow(norm-1.0, aspect='auto', origin='lower', extent=[0, W-1, t[0], t[-1]], vmin=-0.03, vmax=0.03)\n    ax1.set_title(f'AIRS spectrogram (norm-1), {pid}')\n    ax1.set_xlabel('pixel (cropped)'); ax1.set_ylabel('time')\n    plt.colorbar(im, ax=ax1, fraction=0.046, pad=0.04, label='(flux/out)-1')\n    # Panel 2: Depth vs wavelength\n    x = np.arange(W)\n    ax2.plot(x, depth, color='tab:blue', lw=1)\n    ax2.fill_between(x, depth-sigma, depth+sigma, color='tab:blue', alpha=0.2, linewidth=0)\n    ax2.set_title('Depth vs wavelength (arb units)'); ax2.set_xlabel('pixel (cropped)'); ax2.set_ylabel('depth')\n    # Panel 3: Out-of-transit baseline (mean over time)\n    ax3.plot(x, np.nanmean(norm[out_m], axis=0), color='tab:gray', lw=1)\n    ax3.set_title('Out-of-transit mean spectrum'); ax3.set_xlabel('pixel (cropped)'); ax3.set_ylabel('mean flux')\n    # Panel 4: Rolling band time series\n    colors = plt.cm.tab10(np.linspace(0,1,B))\n    for i,(a,b,y) in enumerate(band_series):\n        ax4.plot(t, y/np.nanmedian(y), color=colors[i], lw=0.8, label=f'cols {a}-{b}')\n    ax4.set_title('Band-averaged time series'); ax4.set_xlabel('time'); ax4.set_ylabel('normalized flux'); ax4.legend(ncol=3)\n    plt.tight_layout()\n    fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n    out = fig_dir / f'{pid}_AIRS-CH0_dashboard.png'\n    plt.savefig(out)\n    print('Saved:', out)\nexcept Exception as e:\n    print('AIRS spectral dashboard skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"267edecc","cell_type":"code","source":"# Visualization: Compact quality summary (text report)\ntry:\n    report = []\n    report.append(f\"Env Kaggle={IN_KAGGLE} Input='{INPUT_DIR}' Work='{WORK_DIR}'\")\n    for name in ['train','wavelengths','axis_info','adc_info','train_star','test_star','sample_submission']:\n        df = meta.get(name)\n        shape = tuple(df.shape) if df is not None else None\n        report.append(f\"meta[{name}]: shape={shape}\")\n    # Check existence of split folders\n    for split in ['train','test']:\n        p = INPUT_DIR / split\n        report.append(f\"exists {split}: {p.exists()} path='{p}'\")\n    text = \"\\n\".join(report)\n    print(text)\n    fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n    with open(fig_dir / 'quality_summary.txt', 'w', encoding='utf-8') as f:\n        f.write(text)\n    print('Saved:', fig_dir / 'quality_summary.txt')\nexcept Exception as e:\n    print('Quality summary skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"2119df7a","cell_type":"code","source":"# Premium Viz: FGS light curve with shaded transit and zoom inset\ntry:\n    from mpl_toolkits.axes_grid1.inset_locator import inset_axes, mark_inset\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    base = INPUT_DIR / split\n    planets = [p for p in base.iterdir() if p.is_dir()] if base.exists() else []\n    pid = planets[0].name if planets else 'SYNTH'\n    have_data = bool(planets)\n    if have_data:\n        pth = signal_path(split, pid, 'FGS1', 0)\n        gain, offset = load_adc_info(meta)\n        masters = load_build_calibration(split, pid, 'FGS1', gain, offset)\n        flux_chunks = []\n        for batch in iter_parquet_rows(pth, batch_rows=6000):\n            cal = calibrate_batch_frames(batch, masters, gain, offset)\n            f3 = fgs_unflatten(cal)\n            flux_chunks.append(aperture_photometry(f3, r=load_hparams().get('fgs_aperture_r',5.0)))\n            if sum(len(a) for a in flux_chunks) >= 15000:\n                break\n        flux = np.concatenate(flux_chunks) if flux_chunks else None\n    else:\n        n = 6000\n        t = np.linspace(0, 1, n)\n        transit = 1 - 0.004*(np.abs(t-0.5)<0.06)\n        sys = 1 + 0.003*np.sin(2*np.pi*2.5*t) + 0.001*np.random.randn(n)\n        flux = transit*sys\n    if flux is None or len(flux) < 50:\n        raise RuntimeError('Insufficient FGS data')\n    t = build_time_axis(meta, 'FGS1', len(flux)) if have_data else t\n    # Detrend baseline for cleaner viewing\n    X = build_design_matrix(flux, order=2)\n    _, yhat, _ = generalized_least_squares(flux, X)\n    nf = flux/np.nanmedian(flux)\n    nd = (flux - yhat + np.nanmedian(flux))/np.nanmedian(flux)\n    T = len(flux)\n    t0, t1 = t[int(0.4*T)], t[int(0.6*T)]\n    fig, ax = plt.subplots(figsize=(12,4), constrained_layout=True)\n    ax.plot(t, nf, color='0.7', lw=0.4, label='raw')\n    ax.plot(t, nd, color='tab:blue', lw=0.8, label='detrended')\n    ax.axvspan(t0, t1, color='tab:blue', alpha=0.08, label='in-transit (naive)')\n    ax.set_title(f'FGS1 light curve with shaded transit and zoom ({pid})')\n    ax.set_xlabel('time'); ax.set_ylabel('normalized flux')\n    ax.legend(loc='upper right', frameon=False)\n    # Inset around the minimum in detrended flux (transit-like region)\n    jmin = np.nanargmin(nd)\n    w = max(20, int(0.02*T))\n    x0 = max(t[0], t[jmin]-0.5*(t[w]-t[0]))\n    x1 = min(t[-1], t[jmin]+0.5*(t[w]-t[0]))\n    axins = inset_axes(ax, width=\"35%\", height=\"60%\", loc='lower left', bbox_to_anchor=(0.05,0.05,0.9,0.9), bbox_transform=ax.transAxes, borderpad=0)\n    axins.plot(t, nd, color='tab:blue', lw=0.8)\n    axins.set_xlim(x0, x1)\n    axins.set_ylim(np.nanmin(nd[jmin-w:jmin+w])-0.001, np.nanmax(nd[jmin-w:jmin+w])+0.001)\n    axins.set_xticks([]); axins.set_yticks([])\n    mark_inset(ax, axins, loc1=2, loc2=4, fc=\"none\", ec=\"0.5\")\n    fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n    out = fig_dir / f'{pid}_FGS1_lc_zoom.png'\n    plt.savefig(out)\n    print('Saved:', out)\nexcept Exception as e:\n    print('FGS premium LC visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"454b263b","cell_type":"code","source":"# Premium Viz: AIRS spectrum with smoothing and zoom inset\ntry:\n    from mpl_toolkits.axes_grid1.inset_locator import inset_axes, mark_inset\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    base = INPUT_DIR / split\n    planets = [p for p in base.iterdir() if p.is_dir()] if base.exists() else []\n    pid = planets[0].name if planets else 'SYNTH'\n    have_data = bool(planets)\n    if have_data:\n        pth = signal_path(split, pid, 'AIRS-CH0', 0)\n        gain, offset = load_adc_info(meta)\n        masters = load_build_calibration(split, pid, 'AIRS-CH0', gain, offset)\n        frames = []\n        for batch in iter_parquet_rows(pth, batch_rows=1200):\n            cal = calibrate_batch_frames(batch, masters, gain, offset)\n            f3 = airs_unflatten(cal)\n            f3 = crop_airs_x(f3, 39, 321)\n            frames.append(f3)\n            if sum(x.shape[0] for x in frames) >= 2000:\n                break\n        cube = np.concatenate(frames, axis=0) if frames else None\n    else:\n        # Synthetic cube similar to earlier\n        T, H, W = 1000, 32, 260\n        t = np.linspace(0, 1, T)\n        waves = np.linspace(1.1, 1.9, W)\n        band = np.exp(-0.5*((waves-1.5)/0.12)**2)\n        sys = 1 + 0.01*np.sin(2*np.pi*2*t)[:,None,None]\n        transit = 1 - (0.012*band[None,None,:]) * (np.abs(t-0.5)<0.08)[:,None,None]\n        cube = sys*transit*(1 + 0.001*np.random.randn(T,H,W))\n    if cube is None or cube.shape[0] < 50:\n        raise RuntimeError('Insufficient AIRS data')\n    T = cube.shape[0]\n    spec = cube.sum(axis=1)  # (T, W)\n    i0, i1 = int(0.4*T), int(0.6*T)\n    in_m = np.zeros(T, dtype=bool); in_m[i0:i1] = True\n    out_m = ~in_m\n    out_med = np.nanmedian(spec[out_m], axis=0)\n    out_med[out_med==0] = np.nan\n    norm = spec / out_med[np.newaxis, :]\n    mean_in = np.nanmean(norm[in_m], axis=0)\n    mean_out = np.nanmean(norm[out_m], axis=0)\n    depth = 1.0 - (mean_in/mean_out)\n    x = np.arange(depth.shape[0])\n    # Savitzky-Golay smoothing if scipy available\n    d_smooth = depth\n    try:\n        from scipy.signal import savgol_filter\n        win = max(7, (len(depth)//25)//2*2+1)\n        d_smooth = savgol_filter(depth, window_length=win, polyorder=2, mode='interp')\n    except Exception:\n        pass\n    fig, ax = plt.subplots(figsize=(12,4), constrained_layout=True)\n    ax.plot(x, depth, color='0.7', lw=0.6, label='raw depth')\n    ax.plot(x, d_smooth, color='tab:purple', lw=1.2, label='smoothed')\n    ax.set_title(f'AIRS quick spectrum with smoothing ({pid})')\n    ax.set_xlabel('pixel (cropped)'); ax.set_ylabel('depth (arb)')\n    ax.legend(loc='best', frameon=False)\n    # Inset on lowest (deepest) region\n    j = int(np.nanargmax(d_smooth))\n    w = max(20, len(d_smooth)//10)\n    a = max(0, j-w//2); b = min(len(d_smooth), j+w//2)\n    axins = inset_axes(ax, width=\"35%\", height=\"60%\", loc='upper right')\n    axins.plot(x, depth, color='0.8', lw=0.6)\n    axins.plot(x, d_smooth, color='tab:purple', lw=1.0)\n    axins.set_xlim(a, b)\n    ymin = np.nanmin(d_smooth[a:b]); ymax = np.nanmax(d_smooth[a:b])\n    pad = 0.05*(ymax - ymin + 1e-9)\n    axins.set_ylim(ymin - pad, ymax + pad)\n    axins.set_xticks([]); axins.set_yticks([])\n    mark_inset(ax, axins, loc1=1, loc2=3, fc='none', ec='0.5')\n    fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n    out = fig_dir / f'{pid}_AIRS-CH0_spectrum_zoom.png'\n    plt.savefig(out)\n    print('Saved:', out)\nexcept Exception as e:\n    print('AIRS premium spectrum visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"765ef03b","cell_type":"code","source":"# Premium Viz: FGS centroid hexbin with marginal histograms\ntry:\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    base = INPUT_DIR / split\n    planets = [p for p in base.iterdir() if p.is_dir()] if base.exists() else []\n    pid = planets[0].name if planets else 'SYNTH'\n    have_data = bool(planets)\n    if have_data:\n        pth = signal_path(split, pid, 'FGS1', 0)\n        gain, offset = load_adc_info(meta)\n        masters = load_build_calibration(split, pid, 'FGS1', gain, offset)\n        cx_chunks, cy_chunks = [], []\n        for batch in iter_parquet_rows(pth, batch_rows=5000):\n            cal = calibrate_batch_frames(batch, masters, gain, offset)\n            f3 = fgs_unflatten(cal)\n            cx, cy = compute_centroids(f3)\n            cx_chunks.append(cx); cy_chunks.append(cy)\n            if sum(len(a) for a in cx_chunks) >= 15000:\n                break\n        cx = np.concatenate(cx_chunks) if cx_chunks else None\n        cy = np.concatenate(cy_chunks) if cy_chunks else None\n    else:\n        n = 15000\n        cx = 16 + 0.2*np.sin(np.linspace(0, 25, n)) + 0.1*np.random.randn(n)\n        cy = 16 + 0.2*np.cos(np.linspace(0, 25, n)) + 0.1*np.random.randn(n)\n    if cx is None or len(cx) < 100:\n        raise RuntimeError('Insufficient centroid data')\n    import matplotlib.gridspec as gridspec\n    fig = plt.figure(figsize=(7,6), constrained_layout=True)\n    gs = gridspec.GridSpec(2, 2, width_ratios=[4,1], height_ratios=[1,4], wspace=0.05, hspace=0.05)\n    ax_main = fig.add_subplot(gs[1,0])\n    ax_x = fig.add_subplot(gs[0,0], sharex=ax_main)\n    ax_y = fig.add_subplot(gs[1,1], sharey=ax_main)\n    hb = ax_main.hexbin(cx, cy, gridsize=60, cmap='viridis', mincnt=1)\n    ax_main.set_xlabel('centroid_x'); ax_main.set_ylabel('centroid_y')\n    cb = fig.colorbar(hb, ax=ax_main, fraction=0.046, pad=0.04)\n    cb.set_label('counts')\n    ax_x.hist(cx, bins=60, color='tab:green', alpha=0.8)\n    ax_y.hist(cy, bins=60, orientation='horizontal', color='tab:orange', alpha=0.8)\n    plt.setp(ax_x.get_xticklabels(), visible=False)\n    plt.setp(ax_y.get_yticklabels(), visible=False)\n    ax_x.tick_params(axis='x', which='both', length=0)\n    ax_y.tick_params(axis='y', which='both', length=0)\n    ax_main.set_title(f'FGS centroid density with marginals ({pid})')\n    fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n    out = fig_dir / f'{pid}_FGS1_centroid_hexbin.png'\n    plt.savefig(out)\n    print('Saved:', out)\nexcept Exception as e:\n    print('Centroid hexbin visualization skipped:', e)","metadata":{},"outputs":[],"execution_count":null},{"id":"36c4c59d","cell_type":"code","source":"# Premium Viz: Correlation heatmap (clustered) + target-correlation bars\ntry:\n    # 1) Load a small FGS dataset (or synth fallback)\n    split = 'test' if (INPUT_DIR / 'test').exists() else 'train'\n    base = INPUT_DIR / split\n    planets = [p for p in base.iterdir() if p.is_dir()] if base.exists() else []\n    pid = planets[0].name if planets else 'SYNTH'\n    have_data = bool(planets)\n\n    if have_data:\n        pth = signal_path(split, pid, 'FGS1', 0)\n        gain, offset = load_adc_info(meta)\n        masters = load_build_calibration(split, pid, 'FGS1', gain, offset)\n        flux_chunks, cx_chunks, cy_chunks = [], [], []\n        for batch in iter_parquet_rows(pth, batch_rows=4000):\n            cal = calibrate_batch_frames(batch, masters, gain, offset)\n            f3 = fgs_unflatten(cal)\n            flux_chunks.append(aperture_photometry(f3, r=5.0))\n            cx, cy = compute_centroids(f3)\n            cx_chunks.append(cx); cy_chunks.append(cy)\n            if sum(len(a) for a in flux_chunks) >= 8000:\n                break\n        flux = np.concatenate(flux_chunks) if flux_chunks else None\n        cx = np.concatenate(cx_chunks) if cx_chunks else None\n        cy = np.concatenate(cy_chunks) if cy_chunks else None\n    else:\n        # Synthetic sequence with centroid-driven systematics + shallow transit\n        n = 3000\n        t = np.linspace(0, 1, n)\n        cx = 16 + 0.1*np.sin(2*np.pi*3*t) + 0.02*np.random.randn(n)\n        cy = 16 + 0.1*np.cos(2*np.pi*2*t) + 0.02*np.random.randn(n)\n        sys = 1 + 0.005*t + 0.002*np.sin(2*np.pi*5*t) + 0.001*cx + 0.001*cy\n        transit = 1 - 0.004*(np.abs(t-0.5)<0.06)\n        flux = sys * transit * (1 + 0.001*np.random.randn(n))\n\n    if flux is None or len(flux) < 50:\n        raise RuntimeError('Insufficient data for correlation visualization')\n\n    # 2) Build labeled design matrix (no PLD to keep matrix compact)\n    order = load_hparams().get('poly_order', 2) if 'load_hparams' in globals() else 2\n    X = build_design_matrix(flux, centroids=(cx, cy), order=order)\n    n_cols = X.shape[1]\n    # Labels: [1, t, t^2, ..., cx, cy, cx^2, cy^2, cx*cy]\n    time_labels = ['t'] + [f't^{k}' for k in range(2, order+1)] if order >= 1 else []\n    labels = ['1'] + time_labels + ['cx', 'cy', 'cx^2', 'cy^2', 'cx*cy']\n    labels = labels[:n_cols]  # guard if order/layout differs\n\n    # 3) Correlation matrix with robust handling of NaNs/const cols\n    def safe_corrcoef(M):\n        M = np.asarray(M, dtype=float)\n        M = M - np.nanmean(M, axis=0, keepdims=True)\n        std = np.nanstd(M, axis=0, ddof=0)\n        std[std == 0] = 1.0\n        Mz = M / std\n        C = np.nan_to_num(np.corrcoef(Mz.T), nan=0.0, posinf=0.0, neginf=0.0)\n        C = np.clip(C, -1, 1)\n        return C\n    C = safe_corrcoef(X)\n\n    # 4) Column ordering: cluster by 1-|corr| if scipy available, else by similarity\n    order_idx = np.arange(n_cols)\n    try:\n        from scipy.cluster.hierarchy import linkage, leaves_list\n        from scipy.spatial.distance import squareform\n        D = 1 - np.abs(C)\n        # Ensure proper condensed form input\n        D = np.clip(D, 0, 2)\n        Z = linkage(squareform(D, checks=False), method='average')\n        order_idx = leaves_list(Z)\n    except Exception:\n        # Fallback: sort by total absolute correlation (group similar features)\n        order_idx = np.argsort(-np.sum(np.abs(C), axis=0))\n\n    C_ord = C[np.ix_(order_idx, order_idx)]\n    labels_ord = [labels[i] if i < len(labels) else f'x{i}' for i in order_idx]\n\n    # 5) Target correlation bars (|corr(y, regressor)|) — exclude intercept if present\n    def corr_with_target(y, X):\n        y = np.asarray(y, float)\n        y = y - np.nanmean(y)\n        ys = np.nanstd(y)\n        ys = ys if ys != 0 else 1.0\n        y /= ys\n        out = []\n        for j in range(X.shape[1]):\n            x = X[:, j].astype(float)\n            x = x - np.nanmean(x)\n            xs = np.nanstd(x)\n            xs = xs if xs != 0 else 1.0\n            x /= xs\n            out.append(float(np.nan_to_num(np.mean(x*y), nan=0.0)))\n        return np.array(out)\n    r = np.abs(corr_with_target(flux, X))\n    # Drop intercept for bar chart if present at col 0\n    if labels and labels[0] == '1':\n        r_no_bias = r[1:]\n        labels_no_bias = labels[1:]\n    else:\n        r_no_bias = r\n        labels_no_bias = labels\n    # Order bars by magnitude (top-K)\n    K = min(10, len(labels_no_bias))\n    bar_idx = np.argsort(-r_no_bias)[:K]\n\n    # 6) Plot: heatmap (masked upper triangle) + side bar chart\n    import matplotlib as mpl\n    import matplotlib.pyplot as plt\n    fig = plt.figure(figsize=(12, 7), constrained_layout=True)\n    gs = fig.add_gridspec(1, 2, width_ratios=[3, 1])\n    axH = fig.add_subplot(gs[0, 0])\n    axB = fig.add_subplot(gs[0, 1])\n\n    # Heatmap (prefer seaborn if available)\n    mask = np.triu(np.ones_like(C_ord, dtype=bool), k=1)\n    try:\n        import seaborn as sns\n        sns.heatmap(\n            C_ord,\n            mask=mask,\n            ax=axH,\n            vmin=-1, vmax=1, center=0,\n            cmap='coolwarm',\n            square=True,\n            cbar_kws={'label': 'correlation'},\n            linewidths=0.5, linecolor='white'\n        )\n    except Exception:\n        # Fallback to matplotlib\n        im = axH.imshow(np.where(mask, np.nan, C_ord), vmin=-1, vmax=1, cmap='coolwarm')\n        cb = fig.colorbar(im, ax=axH)\n        cb.set_label('correlation')\n\n    axH.set_xticks(np.arange(len(labels_ord)))\n    axH.set_yticks(np.arange(len(labels_ord)))\n    axH.set_xticklabels(labels_ord, rotation=45, ha='right')\n    axH.set_yticklabels(labels_ord)\n    axH.set_title(f'FGS1 design correlation (clustered) — {pid}')\n\n    # Annotate numbers only for small matrices to avoid clutter\n    if C_ord.shape[0] <= 18:\n        for i in range(C_ord.shape[0]):\n            for j in range(i+1):  # lower triangle incl. diagonal\n                val = C_ord[i, j]\n                axH.text(j, i, f\"{val:.2f}\", ha='center', va='center', fontsize=7,\n                         color=('white' if abs(val) > 0.75 else 'black'))\n\n    # Bar chart: |corr(y, regressor)|\n    axB.barh(range(K), r_no_bias[bar_idx][::-1], color=plt.cm.Blues(np.linspace(0.4, 0.9, K)))\n    axB.set_yticks(range(K))\n    axB.set_yticklabels([labels_no_bias[i] for i in bar_idx][::-1])\n    axB.invert_yaxis()\n    axB.set_xlim(0, 1)\n    axB.set_xlabel('|corr with flux|')\n    axB.set_title('Top drivers')\n\n    # Figure title and save\n    fig.suptitle('Design diagnostics: correlation structure and key drivers', fontsize=13)\n    fig_dir = WORK_DIR / 'figures'; fig_dir.mkdir(parents=True, exist_ok=True)\n    out = fig_dir / f'{pid}_FGS1_corr_annot.png'\n    plt.savefig(out, dpi=200)\n    # Optional SVG for vector clarity\n    try:\n        plt.savefig(out.with_suffix('.svg'))\n    except Exception:\n        pass\n    print('Saved:', out)\nexcept Exception as e:\n    print('Annotated correlation heatmap skipped:', e)","metadata":{},"outputs":[],"execution_count":null}]}