{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":101849,"databundleVersionId":13093295,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13085089,"sourceType":"datasetVersion","datasetId":8287751},{"sourceId":13086430,"sourceType":"datasetVersion","datasetId":8286154},{"sourceId":13160483,"sourceType":"datasetVersion","datasetId":8338925},{"sourceId":589529,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":440860,"modelId":457299}],"dockerImageVersionId":31090,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#import kagglehub\n\n# Download latest version\n#path = kagglehub.dataset_download(\"long2003/ariel-2025-mamba-casualconv1d\")\n\n#print(\"Path to dataset files:\", path)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:28.599041Z","iopub.execute_input":"2025-09-24T18:00:28.599278Z","iopub.status.idle":"2025-09-24T18:00:28.602447Z","shell.execute_reply.started":"2025-09-24T18:00:28.599256Z","shell.execute_reply":"2025-09-24T18:00:28.60183Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"--- Installing necessary packages... ---\")\n# Install required packages quietly|\n!pip install mamba-ssm causal-conv1d einops scikit-learn scipy --quiet \\\n  --no-index \\\n  --find-links=/kaggle/input/ariel-2025-mamba-casualconv1d/\n\n!pip install pqdm --quiet \\\n  --no-index \\\n  --find-links=/kaggle/input/pqdm-dependency/pqdm_package/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:28.603131Z","iopub.execute_input":"2025-09-24T18:00:28.603304Z","iopub.status.idle":"2025-09-24T18:00:34.993095Z","shell.execute_reply.started":"2025-09-24T18:00:28.60329Z","shell.execute_reply":"2025-09-24T18:00:34.992015Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Imports, Installations, and Configuration","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 2. Imports and Configuration ---\nimport os\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport pandas as pd\nimport numpy as np\nimport torch.nn.functional as F\nfrom tqdm.notebook import tqdm\nfrom torch.utils.data import DataLoader, Dataset, Subset\nfrom mamba_ssm import Mamba\nfrom sklearn.preprocessing import StandardScaler\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom sklearn.model_selection import KFold\nimport pandas as pd\nimport numpy as np\nfrom tqdm import tqdm\nfrom pqdm.threads import pqdm\nfrom scipy.signal import savgol_filter\nfrom astropy.stats import sigma_clip\nimport time\nimport sys\nimport glob\nimport pickle\nimport gc\nimport pywt\nfrom statsmodels.robust import mad\nfrom dataclasses import dataclass, field\nfrom typing import Dict, Tuple\n\n\n# --- GPU Device Setup ---\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {DEVICE}\")\n\n@dataclass\nclass Config:\n    \"\"\"Configuration class for the advanced Ariel data processing pipeline.\"\"\"\n    # Data paths\n    DATA_PATH: str = '/kaggle/input/ariel-data-challenge-2025'\n    PROCESSED_DATA_PATH: str = '/kaggle/input/d/longddang/ariel-processed-data/'\n    OUTPUT_PATH: str = '/kaggle/working/'\n    CHECKPOINT_FILE: str = 'processed_data_checkpoint.npz'\n    TABULAR_SCALER_FILE: str = '/kaggle/input/d/longddang/ariel-processed-data/tabular_scaler.pkl'\n    TS_SCALER_FILE: str = '/kaggle/input/d/longddang/ariel-processed-data/ts_scaler.pkl'\n    \n    # Dataset and processing\n    DATASET_TYPE: str = 'train'  # Process the training data\n    N_JOBS: int = os.cpu_count()  # Use all available CPU cores\n    BATCH_SIZE_PROCESSING: int = 100  # Process N planets at a time to manage memory\n    \n    # Sensor configurations\n    SENSOR_CONFIG: Dict = field(default_factory=lambda: {\n        \"AIRS-CH0\": {\n            \"raw_shape\": [11250, 32, 356],\n            \"linear_corr_shape\": (6, 32, 356),\n            \"dt_pattern\": (0.1, 4.5),\n            \"binning\": 30,\n            \"roi_y\": (10, 22),\n            \"cut_inf\": 39,\n            \"cut_sup\": 321,\n        },\n        \"FGS1\": {\n            \"raw_shape\": [135000, 32, 32],\n            \"linear_corr_shape\": (6, 32, 32),\n            \"dt_pattern\": (0.1, 0.1),\n            \"binning\": 30 * 12,  # Match AIRS binning time\n            \"roi_y\": (10, 22),\n            \"roi_x\": (10, 22)\n        }\n    })\n    \n    # Model & Training Parameters\n    MDN_COMPONENTS: int = 5  # Number of Gaussian mixtures\n    N_SPLITS: int = 10\n    BATCH_SIZE: int = 64  # Batch size for training\n    NUM_EPOCHS: int = 100\n    EARLY_STOPPING_PATIENCE: int = 15\n    LEARNING_RATE: float = 2e-5\n    GRAD_CLIP_VALUE: float = 1.0\n    DROPOUT_RATE: float = 0.1\n    \n    # Mamba Architecture\n    MAMBA_D_MODEL: int = 128\n    MAMBA_NUM_LAYERS: int = 4\n    MAMBA_D_STATE: int = 16\n    MAMBA_D_CONV: int = 4\n    MAMBA_EXPAND: int = 2\n    TABULAR_HIDDEN_DIM: int = 128\n    HEAD_HIDDEN_DIM: int = 256","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:34.994412Z","iopub.execute_input":"2025-09-24T18:00:34.994719Z","iopub.status.idle":"2025-09-24T18:00:35.008202Z","shell.execute_reply.started":"2025-09-24T18:00:34.99469Z","shell.execute_reply":"2025-09-24T18:00:35.007557Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Signal Processor for TEST data only","metadata":{}},{"cell_type":"code","source":"# --- 3. Advanced Signal Processing Class ---\nclass Signal_Processor:\n    \"\"\"Processes sensor signals using a full calibration and extraction pipeline.\"\"\"\n    def __init__(self, config: Config):\n        self.cfg = config\n        self.adc_info = pd.read_csv(f\"{self.cfg.DATA_PATH}/adc_info.csv\")\n        self.meta_info = pd.read_csv(f'{self.cfg.DATA_PATH}/{self.cfg.DATASET_TYPE}_star_info.csv')\n        self.planet_ids = self.meta_info['planet_id'].unique()\n\n    def _apply_linear_corr(self, linear_corr: np.ndarray, signal: np.ndarray) -> np.ndarray:\n        \"\"\"Apply linearity correction to the signal.\"\"\"\n        coeffs = np.flip(linear_corr, axis=0)\n        x = signal.astype(np.float64, copy=False)\n        out = np.empty_like(x, dtype=np.float64)\n        out[...] = coeffs[0]\n        for k in range(1, coeffs.shape[0]):\n            np.multiply(out, x, out=out)\n            out += coeffs[k]\n        return out.astype(signal.dtype, copy=False)\n\n    def _calibrate_single_signal(self, planet_id: int, sensor: str) -> np.ndarray:\n        \"\"\"Calibrate a single sensor signal for a planet.\"\"\"\n        sensor_cfg = self.cfg.SENSOR_CONFIG[sensor]\n        base_path = f\"{self.cfg.DATA_PATH}/{self.cfg.DATASET_TYPE}/{planet_id}\"\n        \n        # Load all required files\n        signal = pd.read_parquet(f\"{base_path}/{sensor}_signal_0.parquet\").to_numpy()\n        dark = pd.read_parquet(f\"{base_path}/{sensor}_calibration_0/dark.parquet\").to_numpy()\n        dead = pd.read_parquet(f\"{base_path}/{sensor}_calibration_0/dead.parquet\").to_numpy()\n        flat = pd.read_parquet(f\"{base_path}/{sensor}_calibration_0/flat.parquet\").to_numpy()\n        linear_corr = pd.read_parquet(f\"{base_path}/{sensor}_calibration_0/linear_corr.parquet\").values.astype(np.float64).reshape(sensor_cfg[\"linear_corr_shape\"])\n\n        # Reshape & apply ADC correction\n        signal = signal.reshape(sensor_cfg[\"raw_shape\"])\n        gain = self.adc_info[f\"{sensor}_adc_gain\"].iloc[0]\n        offset = self.adc_info[f\"{sensor}_adc_offset\"].iloc[0]\n        signal = signal / gain + offset\n\n        # Apply sensor-specific cropping for AIRS-CH0\n        if sensor == \"AIRS-CH0\":\n            cut_inf, cut_sup = sensor_cfg[\"cut_inf\"], sensor_cfg[\"cut_sup\"]\n            signal = signal[:, :, cut_inf:cut_sup]\n            linear_corr = linear_corr[:, :, cut_inf:cut_sup]\n            dark = dark[:, cut_inf:cut_sup]\n            dead = dead[:, cut_inf:cut_sup]\n            flat = flat[:, cut_inf:cut_sup]\n\n        # Clamp to non-negative before linearity correction\n        np.maximum(signal, 0, out=signal)\n\n        # Apply linearity correction to the region of interest\n        y0, y1 = sensor_cfg[\"roi_y\"]\n        if sensor == \"FGS1\":\n            x0, x1 = sensor_cfg[\"roi_x\"]\n            roi_slice = (slice(None), slice(y0, y1), slice(x0, x1))\n            signal[roi_slice] = self._apply_linear_corr(linear_corr[:, y0:y1, x0:x1], signal[roi_slice])\n        elif sensor == \"AIRS-CH0\":\n            roi_slice = (slice(None), slice(y0, y1), slice(None))\n            signal[roi_slice] = self._apply_linear_corr(linear_corr[:, y0:y1, :], signal[roi_slice])\n\n        # Dark subtraction\n        base_dt, increment = sensor_cfg[\"dt_pattern\"]\n        signal[::2] -= dark * base_dt\n        signal[1::2] -= dark * (base_dt + increment)\n\n        # Flat field correction (avoiding dead pixels)\n        flat[dead] = np.nan\n        signal /= flat\n        \n        return signal\n\n    def _wavelet_denoise(self, signal: np.ndarray, wavelet='db4', level=1) -> np.ndarray:\n        \"\"\"Denoise a signal using wavelet transform.\"\"\"\n        coeff = pywt.wavedec(signal, wavelet, mode=\"per\")\n        # Ensure mad is calculated robustly\n        sigma = mad(coeff[-level], center=np.median)\n        # Universal Threshold\n        uthresh = sigma * np.sqrt(2 * np.log(len(signal)))\n        coeff[1:] = (pywt.threshold(i, value=uthresh, mode='soft') for i in coeff[1:])\n        return pywt.waverec(coeff, wavelet, mode='per')\n\n    def _extract_binned_cds(self, calibrated_signal: np.ndarray, sensor: str) -> np.ndarray:\n        \"\"\"Extract CDS signal from ROI, denoise, bin it, and apply weighting.\"\"\"\n        sensor_cfg = self.cfg.SENSOR_CONFIG[sensor]\n        y0, y1 = sensor_cfg[\"roi_y\"]\n        \n        if sensor == \"FGS1\":\n            x0, x1 = sensor_cfg[\"roi_x\"]\n            signal_roi = calibrated_signal[:, y0:y1, x0:x1]\n            mean_signal = np.nanmean(signal_roi, axis=(1, 2))\n        else: # AIRS-CH0\n            signal_roi = calibrated_signal[:, y0:y1, :]\n            mean_signal = np.nanmean(signal_roi, axis=1)\n\n        # Correlated Double Sampling (CDS)\n        cds_signal = mean_signal[1::2] - mean_signal[0::2]\n        \n        # Apply Wavelet Denoising\n        if cds_signal.ndim > 1: # For AIRS data (Time, Wavelengths)\n            denoised_signal = np.apply_along_axis(self._wavelet_denoise, 0, cds_signal)\n        else: # For FGS1 data (Time,)\n            denoised_signal = self._wavelet_denoise(cds_signal)\n\n        # Bin the denoised signal\n        binning = sensor_cfg[\"binning\"]\n        n_bins = denoised_signal.shape[0] // binning\n        binned_shape = (n_bins,) + denoised_signal.shape[1:]\n        binned = np.zeros(binned_shape)\n        for j in range(n_bins):\n            binned[j] = np.nanmean(denoised_signal[j*binning : (j+1)*binning], axis=0)\n\n        # Reshape FGS1 to be compatible\n        if sensor == \"FGS1\":\n            binned = binned.reshape((binned.shape[0], 1))\n\n        return np.nan_to_num(binned)\n\n    def process_planet(self, planet_id: int) -> Tuple[int, np.ndarray | None]:\n        \"\"\"Wrapper for parallel processing of a single planet.\"\"\"\n        try:\n            calibrated_fgs1 = self._calibrate_single_signal(planet_id, \"FGS1\")\n            calibrated_airs = self._calibrate_single_signal(planet_id, \"AIRS-CH0\")\n            \n            processed_fgs1 = self._extract_binned_cds(calibrated_fgs1, \"FGS1\")\n            processed_airs = self._extract_binned_cds(calibrated_airs, \"AIRS-CH0\")\n\n            # Concatenate along the feature axis\n            combined_signal = np.concatenate([processed_fgs1, processed_airs], axis=1)\n            return planet_id, combined_signal\n        \n        except Exception as e:\n            print(e)\n            return planet_id, None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:35.008904Z","iopub.execute_input":"2025-09-24T18:00:35.009162Z","iopub.status.idle":"2025-09-24T18:00:35.332058Z","shell.execute_reply.started":"2025-09-24T18:00:35.009136Z","shell.execute_reply":"2025-09-24T18:00:35.331367Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## PyTorch Dataset & MDN Mamba Model","metadata":{}},{"cell_type":"code","source":"# --- 2. PyTorch Dataset & MDN Mamba Model ---\nclass ArielDataset(Dataset):\n    def __init__(self, ts_features, tabular_features, targets=None, is_train=True):\n        self.ts_features = ts_features.astype(np.float32)\n        self.tabular_features = tabular_features.astype(np.float32)\n        self.is_train = is_train\n        if is_train:\n            self.targets = targets.astype(np.float32)\n\n    def __len__(self):\n        return len(self.ts_features)\n\n    def __getitem__(self, idx):\n        if self.is_train:\n            return torch.tensor(self.ts_features[idx]), torch.tensor(self.tabular_features[idx]), torch.tensor(self.targets[idx])\n        else:\n            return torch.tensor(self.ts_features[idx]), torch.tensor(self.tabular_features[idx])\n\n\nclass MambaEncoder(nn.Module):\n    def __init__(self, input_dim, d_model, num_layers, **kwargs):\n        super().__init__()\n        self.embedding = nn.Linear(input_dim, d_model)\n        self.mamba_layers = nn.ModuleList([Mamba(d_model=d_model, **kwargs) for _ in range(num_layers)])\n        self.norm = nn.LayerNorm(d_model)\n        \n    def forward(self, x):\n        x = self.embedding(x)\n        for layer in self.mamba_layers:\n            x = self.norm(x + layer(x))\n        return torch.mean(x, dim=1)\n\n\nclass MDNHead(nn.Module):\n    def __init__(self, latent_dim, output_dim, hidden_dim, dropout_rate, n_components):\n        super().__init__()\n        self.n_components = n_components\n        self.output_dim = output_dim\n        # Each component needs pi (weight), mu (mean), and sigma (std dev)\n        self.output_multiplier = n_components * 3 \n        \n        self.net = nn.Sequential(\n            nn.Linear(latent_dim, hidden_dim), nn.SiLU(), nn.Dropout(dropout_rate),\n            nn.Linear(hidden_dim, hidden_dim), nn.SiLU(), nn.Dropout(dropout_rate),\n            nn.Linear(hidden_dim, output_dim * self.output_multiplier)\n        )\n        \n    def forward(self, latent):\n        params = self.net(latent)\n        # Reshape to (batch, output_dim, n_components, 3)\n        params = params.view(latent.size(0), self.output_dim, self.n_components, 3)\n        \n        # Split into pi, mu, sigma\n        pi_logits = params[..., 0]\n        mu = params[..., 1]\n        sigma_logits = params[..., 2]\n        \n        # Apply activations\n        pi = F.softmax(pi_logits, dim=2)\n        # Use softplus to ensure sigma is positive and stable\n        sigma = F.softplus(sigma_logits) + 1e-6 \n        \n        return pi, mu, sigma\n\nclass MambaMDNModel(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        self.ts_encoder = MambaEncoder(input_dim=cfg.TS_INPUT_DIM, \n                                       d_model=cfg.MAMBA_D_MODEL, \n                                       num_layers=cfg.MAMBA_NUM_LAYERS, \n                                       d_state=cfg.MAMBA_D_STATE, \n                                       d_conv=cfg.MAMBA_D_CONV, \n                                       expand=cfg.MAMBA_EXPAND)\n        \n        self.tabular_encoder = nn.Sequential(\n            nn.Linear(cfg.TABULAR_INPUT_DIM, cfg.TABULAR_HIDDEN_DIM), nn.ReLU(),\n            nn.Linear(cfg.TABULAR_HIDDEN_DIM, cfg.TABULAR_HIDDEN_DIM // 2), nn.ReLU()\n        )\n        \n        combined_dim = cfg.MAMBA_D_MODEL + (cfg.TABULAR_HIDDEN_DIM // 2)\n        self.head = MDNHead(latent_dim=combined_dim, output_dim=cfg.TARGET_SPECTRUM_DIM, \n                            hidden_dim=cfg.HEAD_HIDDEN_DIM, dropout_rate=cfg.DROPOUT_RATE, \n                            n_components=cfg.MDN_COMPONENTS)\n\n    def forward(self, ts_input, tabular_input):\n        ts_latent = self.ts_encoder(ts_input)\n        tabular_latent = self.tabular_encoder(tabular_input)\n        combined_latent = torch.cat([ts_latent, tabular_latent], dim=1)\n        return self.head(combined_latent)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:35.332879Z","iopub.execute_input":"2025-09-24T18:00:35.333139Z","iopub.status.idle":"2025-09-24T18:00:35.346211Z","shell.execute_reply.started":"2025-09-24T18:00:35.333114Z","shell.execute_reply":"2025-09-24T18:00:35.345547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 3. Training Utilities ---\ndef mdn_loss(pi, mu, sigma, target):\n    target = target.unsqueeze(2) # Make target shape broadcastable with mixture params\n    \n    # Calculate log probability of target under each Gaussian component\n    m = torch.distributions.Normal(loc=mu, scale=sigma)\n    log_prob = m.log_prob(target)\n    \n    # Weight the log probabilities by the mixture weights (in log space)\n    log_pi = torch.log(pi)\n    weighted_log_prob = log_prob + log_pi\n    \n    # Sum probabilities in a numerically stable way\n    log_sum = torch.logsumexp(weighted_log_prob, dim=2)\n    \n    # Return the negative mean log-likelihood\n    return -log_sum.mean()\n\n\ndef train_model(model, train_loader, val_loader, optimizer, \n                scheduler, num_epochs, grad_clip_value, fold, patience):\n    scaler = torch.amp.GradScaler(device='cuda', enabled=torch.cuda.is_available())\n    best_val_loss = np.inf\n    epochs_no_improve = 0\n    \n    for epoch in range(num_epochs):\n        model.train()\n        total_train_loss = 0\n        pbar = tqdm(train_loader, desc=f\"Fold {fold+1} Epoch {epoch+1}/{num_epochs}\", \n                    leave=False)\n        for ts_inputs, tab_inputs, targets in pbar:\n            ts_inputs, tab_inputs, targets = ts_inputs.to(DEVICE), \\\n                                             tab_inputs.to(DEVICE), \\\n                                             targets.to(DEVICE)\n            optimizer.zero_grad(set_to_none=True)\n            with torch.amp.autocast(device_type='cuda', dtype=torch.float32, enabled=torch.cuda.is_available()):\n                pi, mu, sigma = model(ts_inputs, tab_inputs)\n                loss = mdn_loss(pi, mu, sigma, targets)\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_value)\n            scaler.step(optimizer)\n            scaler.update()\n            total_train_loss += loss.item()\n        avg_train_loss = total_train_loss / len(train_loader)\n        model.eval()\n        total_val_loss = 0\n        with torch.no_grad():\n            for ts_inputs, tab_inputs, targets in val_loader:\n                ts_inputs, tab_inputs, targets = ts_inputs.to(DEVICE), tab_inputs.to(DEVICE), targets.to(DEVICE)\n                with torch.amp.autocast(device_type='cuda', dtype=torch.float32, enabled=torch.cuda.is_available()):\n                    pi, mu, sigma = model(ts_inputs, tab_inputs)\n                    total_val_loss += mdn_loss(pi, mu, sigma, targets).item()\n        avg_val_loss = total_val_loss / len(val_loader)\n        scheduler.step(avg_val_loss)\n        print(f\"Fold {fold+1} Epoch {epoch+1} | Train Loss: {avg_train_loss:.6f} | Val Loss: {avg_val_loss:.6f}\")\n        if avg_val_loss < best_val_loss - 1e-6:\n            best_val_loss = avg_val_loss\n            state_dict_to_save = model._orig_mod.state_dict() if hasattr(model, '_orig_mod') else model.state_dict()\n            torch.save(state_dict_to_save, f\"mamba_mdn_model_fold_{fold+1}.pth\")\n            epochs_no_improve = 0\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= patience:\n                print(f\"    -> Early stopping triggered for Fold {fold+1}.\")\n                break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:35.347023Z","iopub.execute_input":"2025-09-24T18:00:35.347249Z","iopub.status.idle":"2025-09-24T18:00:35.361899Z","shell.execute_reply.started":"2025-09-24T18:00:35.347225Z","shell.execute_reply":"2025-09-24T18:00:35.361267Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"# --- STAGE 1: Load Pre-processed Data ---\nprint(\"--- STAGE 1: Loading Pre-processed Data from Dataset ---\")\nCONFIG = Config()\n\ncheckpoint_path = os.path.join(CONFIG.PROCESSED_DATA_PATH, \n                               CONFIG.CHECKPOINT_FILE)\nif not os.path.exists(checkpoint_path):\n    print(f\"🔴 FATAL ERROR: Checkpoint file not found at '{checkpoint_path}'.\")\n    sys.exit(\"Please ensure the 'Ariel Processed Data' dataset is attached.\")\n    \ncheckpoint = np.load(checkpoint_path)\nprocessed_train_data = checkpoint['processed_train_data']\ntrain_targets = checkpoint['train_targets']\nscaled_tabular_train = checkpoint['tabular_features']\n\nprint(\"--- Data loaded successfully from checkpoint. ---\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:35.362607Z","iopub.execute_input":"2025-09-24T18:00:35.362791Z","iopub.status.idle":"2025-09-24T18:00:38.351682Z","shell.execute_reply.started":"2025-09-24T18:00:35.362777Z","shell.execute_reply":"2025-09-24T18:00:38.350992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# n_samples, n_timesteps, n_features = processed_train_data.shape\n# ts_data_reshaped = processed_train_data.reshape(-1, n_features)\n# ts_scaler = StandardScaler()\n# ts_data_scaled_reshaped = ts_scaler.fit_transform(ts_data_reshaped)\n# scaled_train_data = ts_data_scaled_reshaped.reshape(n_samples, n_timesteps, n_features)\n\n# # # with open(CONFIG['TS_SCALER_FILE'], 'wb') as f: pickle.dump(ts_scaler, f)\n# tabular_scaler_path = os.path.join(CONFIG.PROCESSED_DATA_PATH, CONFIG.TABULAR_SCALER_FILE)\n# with open(tabular_scaler_path, 'rb') as f: \n#     tabular_scaler = pickle.load(f)\n\n# # with open(CONFIG['TABULAR_SCALER_FILE'], 'wb') as f: pickle.dump(tabular_scaler, f)\n# print(\"--- Scalers saved to output. ---\")\n\n# CONFIG.TS_INPUT_DIM = scaled_train_data.shape[2]\n# CONFIG.TABULAR_INPUT_DIM = scaled_tabular_train.shape[1]\n# CONFIG.TARGET_SPECTRUM_DIM = train_targets.shape[1]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:38.352412Z","iopub.execute_input":"2025-09-24T18:00:38.352619Z","iopub.status.idle":"2025-09-24T18:00:38.35627Z","shell.execute_reply.started":"2025-09-24T18:00:38.352605Z","shell.execute_reply":"2025-09-24T18:00:38.355656Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(\"\\n--- STAGE 2: Mamba MDN Model Training ---\")\n\n# full_dataset = ArielDataset(scaled_train_data, scaled_tabular_train, train_targets, is_train=True)\n# kfold = KFold(n_splits=CONFIG.N_SPLITS, shuffle=True, random_state=42)\n\n# for fold, (train_ids, val_ids) in enumerate(kfold.split(full_dataset)):\n#     print(f\"\\n--- Starting Fold {fold+1}/{CONFIG.N_SPLITS} ---\")\n    \n#     train_loader = DataLoader(Subset(full_dataset, train_ids), \n#                               batch_size=CONFIG.BATCH_SIZE, \n#                               shuffle=True, num_workers=2, \n#                               pin_memory=True)\n#     val_loader = DataLoader(Subset(full_dataset, val_ids), \n#                             batch_size=CONFIG.BATCH_SIZE, \n#                             shuffle=False, num_workers=2, \n#                             pin_memory=True)\n    \n#     model = MambaMDNModel(CONFIG).to(DEVICE)\n#     if torch.__version__ >= \"2.0.0\": model = torch.compile(model)\n        \n#     optimizer = optim.AdamW(model.parameters(), lr=CONFIG.LEARNING_RATE)\n#     scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=10)\n#     train_model(model, train_loader, val_loader, optimizer, scheduler, \n#                 CONFIG.NUM_EPOCHS, CONFIG.GRAD_CLIP_VALUE, \n#                 fold, CONFIG.EARLY_STOPPING_PATIENCE)\n#     del model, train_loader, val_loader; gc.collect(); torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:38.357171Z","iopub.execute_input":"2025-09-24T18:00:38.357723Z","iopub.status.idle":"2025-09-24T18:00:38.371084Z","shell.execute_reply.started":"2025-09-24T18:00:38.357698Z","shell.execute_reply":"2025-09-24T18:00:38.370418Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Utilities","metadata":{}},{"cell_type":"markdown","source":"#### Load Preprocessing Data","metadata":{}},{"cell_type":"code","source":"\nprint(\"\\n--- Scaling Time-Series Data ---\")\nn_samples, n_timesteps, n_features = processed_train_data.shape\nts_data_reshaped = processed_train_data.reshape(-1, n_features)\nts_scaler = StandardScaler()\nts_data_scaled_reshaped = ts_scaler.fit_transform(ts_data_reshaped)\nscaled_train_data = ts_data_scaled_reshaped.reshape(n_samples, n_timesteps, n_features)\n\n# with open(CONFIG['TS_SCALER_FILE'], 'wb') as f: pickle.dump(ts_scaler, f)\ntabular_scaler_path = os.path.join(CONFIG.PROCESSED_DATA_PATH, CONFIG.TABULAR_SCALER_FILE)\nwith open(tabular_scaler_path, 'rb') as f: tabular_scaler = pickle.load(f)\n# with open(CONFIG['TABULAR_SCALER_FILE'], 'wb') as f: pickle.dump(tabular_scaler, f)\nprint(\"--- Scalers saved to output. ---\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:38.374043Z","iopub.execute_input":"2025-09-24T18:00:38.374247Z","iopub.status.idle":"2025-09-24T18:00:39.283919Z","shell.execute_reply.started":"2025-09-24T18:00:38.374233Z","shell.execute_reply":"2025-09-24T18:00:39.283225Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CONFIG['TS_INPUT_DIM'] = scaled_train_data.shape[2]\n# CONFIG['TABULAR_INPUT_DIM'] = scaled_tabular_train.shape[1]\n# CONFIG['TARGET_SPECTRUM_DIM'] = train_targets.shape[1]\n\n# --- STAGE 2: MAMBA MDN MODEL TRAINING ---\n# print(\"\\n--- STAGE 2: Mamba MDN Model Training ---\")\n\n\n# full_dataset = ArielDataset(scaled_train_data, scaled_tabular_train, train_targets, is_train=True)\n# kfold = KFold(n_splits=CONFIG['N_SPLITS'], shuffle=True, random_state=42)\n\n# for fold, (train_ids, val_ids) in enumerate(kfold.split(full_dataset)):\n#     print(f\"\\n--- Starting Fold {fold+1}/{CONFIG['N_SPLITS']} ---\")\n    \n#     train_loader = DataLoader(Subset(full_dataset, train_ids), batch_size=CONFIG['BATCH_SIZE'], shuffle=True, num_workers=2, pin_memory=True)\n#     val_loader = DataLoader(Subset(full_dataset, val_ids), batch_size=CONFIG['BATCH_SIZE'], shuffle=False, num_workers=2, pin_memory=True)\n    \n#     model = MambaMDNModel(CONFIG).to(DEVICE)\n#     if torch.__version__ >= \"2.0.0\": model = torch.compile(model)\n        \n#     optimizer = optim.AdamW(model.parameters(), lr=CONFIG['LEARNING_RATE'])\n#     scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=10)\n#     train_model(model, train_loader, val_loader, optimizer, scheduler, CONFIG['NUM_EPOCHS'], CONFIG['GRAD_CLIP_VALUE'], fold, CONFIG['EARLY_STOPPING_PATIENCE'])\n#     del model, train_loader, val_loader; gc.collect(); torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:39.284495Z","iopub.execute_input":"2025-09-24T18:00:39.284727Z","iopub.status.idle":"2025-09-24T18:00:39.28914Z","shell.execute_reply.started":"2025-09-24T18:00:39.284707Z","shell.execute_reply":"2025-09-24T18:00:39.288408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weight_path = \"/kaggle/input/ariel_mdn_model/pytorch/updated_signal/1\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:39.289788Z","iopub.execute_input":"2025-09-24T18:00:39.290057Z","iopub.status.idle":"2025-09-24T18:00:39.305261Z","shell.execute_reply.started":"2025-09-24T18:00:39.290036Z","shell.execute_reply":"2025-09-24T18:00:39.304636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- STAGE 4: SUBMISSION GENERATION ---\nprint(\"\\n--- STAGE 4: Submission Generation ---\")\nCONFIG = Config()\nCONFIG.DATASET_TYPE = 'test'\nCONFIG.TS_INPUT_DIM = scaled_train_data.shape[2]\nCONFIG.TABULAR_INPUT_DIM = scaled_tabular_train.shape[1]\nCONFIG.TARGET_SPECTRUM_DIM = train_targets.shape[1]\n\n\ndf_meta_test = pd.read_csv(os.path.join(CONFIG.DATA_PATH, 'test_star_info.csv'))\ndf_adc = pd.read_csv(os.path.join(CONFIG.DATA_PATH, 'adc_info.csv'))\nadc_info_row = df_adc.iloc[0]\n\ntest_data_path = os.path.join(CONFIG.DATA_PATH, 'test')\n\ndiscovered_planet_dirs = [d for d in os.listdir(test_data_path) if os.path.isdir(os.path.join(test_data_path, d)) and d.isdigit()]\ntest_planet_ids = [int(pid) for pid in discovered_planet_dirs]\nprint(f\"Discovered {len(test_planet_ids)} planet directories in the test set.\")\n\ntest_processor = Signal_Processor(CONFIG)\ntest_results = pqdm(test_planet_ids, test_processor.process_planet, n_jobs=os.cpu_count())\nprocessed_test_map = {pid: data for pid, data in test_results if data is not None}\nsuccessful_test_ids = list(processed_test_map.keys())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:39.306104Z","iopub.execute_input":"2025-09-24T18:00:39.30632Z","iopub.status.idle":"2025-09-24T18:00:44.472383Z","shell.execute_reply.started":"2025-09-24T18:00:39.306298Z","shell.execute_reply":"2025-09-24T18:00:44.471465Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if successful_test_ids:\n    processed_test_data = np.array([processed_test_map[pid] for pid in successful_test_ids])\n    \n    with open(CONFIG.TABULAR_SCALER_FILE, 'rb') as f: \n        tabular_scaler = pickle.load(f)\n    with open(CONFIG.TS_SCALER_FILE, 'rb') as f: \n        ts_scaler = pickle.load(f)\n        \n    n_samples, n_timesteps, n_features = processed_test_data.shape\n    ts_test_reshaped = processed_test_data.reshape(-1, n_features)\n    ts_test_scaled_reshaped = ts_scaler.transform(ts_test_reshaped)\n    scaled_test_ts_data = ts_test_scaled_reshaped.reshape(n_samples, n_timesteps, n_features)\n        \n    tab_cols = ['Rs', 'Ms', 'Ts', 'Mp', 'e', 'P', 'sma', 'i']\n    df_meta_test_filtered = df_meta_test[df_meta_test.planet_id.isin(successful_test_ids)].set_index('planet_id').loc[successful_test_ids]\n    tabular_features_test = df_meta_test_filtered[tab_cols].fillna(0)\n    scaled_tabular_test = tabular_scaler.transform(tabular_features_test)\n    \n    test_dataset = ArielDataset(scaled_test_ts_data, scaled_tabular_test, is_train=False)\n    test_loader = DataLoader(test_dataset, batch_size=CONFIG.BATCH_SIZE, shuffle=False, num_workers=2)\n    \n    all_fold_means, all_fold_stds = [], []\n    for fold in range(CONFIG.N_SPLITS):\n        model = MambaMDNModel(CONFIG).to(DEVICE)\n        model.load_state_dict(torch.load(f\"{weight_path}/mamba_mdn_model_fold_{fold+1}.pth\"))\n        if torch.__version__ >= \"2.0.0\": model = torch.compile(model)\n        model.eval()\n        \n        fold_pis, fold_mus, fold_sigmas = [], [], []\n        with torch.no_grad():\n            for ts_inputs, tab_inputs in tqdm(test_loader, desc=f\"Predicting Fold {fold+1}\"):\n                ts_inputs, tab_inputs = ts_inputs.to(DEVICE), tab_inputs.to(DEVICE)\n                with torch.amp.autocast(device_type='cuda', dtype=torch.float32, enabled=torch.cuda.is_available()):\n                    pi, mu, sigma = model(ts_inputs, tab_inputs)\n                fold_pis.append(pi.cpu().numpy())\n                fold_mus.append(mu.cpu().numpy())\n                fold_sigmas.append(sigma.cpu().numpy())\n        \n        pi = np.concatenate(fold_pis); \n        mu = np.concatenate(fold_mus); \n        sigma = np.concatenate(fold_sigmas)\n        mean = np.sum(pi * mu, axis=2)\n        var = np.sum(pi * (np.square(mu) + np.square(sigma)), axis=2) - np.square(mean)\n        std = np.sqrt(np.maximum(var, 1e-9))\n        \n        all_fold_means.append(mean)\n        all_fold_stds.append(std)\n\n    final_mean_predictions = np.mean(all_fold_means, axis=0)\n    final_std_predictions = np.mean(all_fold_stds, axis=0) # Simple averaging for std\n    \n    N_OUTPUTS = final_mean_predictions.shape[1]\n    \n    sample_submission = pd.read_csv(\"/kaggle/input/ariel-data-challenge-2025/sample_submission.csv\",\n                                   index_col =\"planet_id\")\n    planet_ids = sample_submission.index\n    \n    oof_mean = np.clip(final_mean_predictions, 0, None)\n    oof_std = np.clip(final_std_predictions, 1e-6, 0.1)\n\n        \n    submission_df = pd.DataFrame(\n        np.concatenate([oof_mean, oof_std], axis = 1),\n        columns = sample_submission.columns,\n        index = planet_ids\n    )    \n    submission_df.to_csv(\"submission.csv\")\n    print(\"\\nSubmission file 'submission.csv' created successfully!\")\n    print(submission_df.head())\nelse:\n    print(\"No test planets processed. Creating submission from sample.\")\n    sample_sub = pd.read_csv(os.path.join(CONFIG.DATA_PATH, 'sample_submission.csv'))\n    sample_sub.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:44.473446Z","iopub.execute_input":"2025-09-24T18:00:44.47426Z","iopub.status.idle":"2025-09-24T18:00:46.562391Z","shell.execute_reply.started":"2025-09-24T18:00:44.47423Z","shell.execute_reply":"2025-09-24T18:00:46.561284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.read_csv(\"submission.csv\")\nsubmission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-24T18:00:46.563635Z","iopub.execute_input":"2025-09-24T18:00:46.563959Z","iopub.status.idle":"2025-09-24T18:00:46.593414Z","shell.execute_reply.started":"2025-09-24T18:00:46.563904Z","shell.execute_reply":"2025-09-24T18:00:46.592781Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}