{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":70367,"databundleVersionId":9188054,"sourceType":"competition"},{"sourceId":9110664,"sourceType":"datasetVersion","datasetId":5498833}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.layers import (Input, Conv1D, MaxPooling1D, Flatten, Dense,\n                                     Dropout, BatchNormalization, Conv2D, MaxPooling2D,\n                                     LSTM, Concatenate, Reshape)\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.losses import MeanAbsoluteError\nfrom tensorflow.keras.callbacks import LearningRateScheduler, ModelCheckpoint\nimport pandas as pd\nimport random\nimport os","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-20T05:04:56.317212Z","iopub.execute_input":"2025-04-20T05:04:56.317566Z","iopub.status.idle":"2025-04-20T05:04:56.322959Z","shell.execute_reply.started":"2025-04-20T05:04:56.317514Z","shell.execute_reply":"2025-04-20T05:04:56.3221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set paths to your data\nDATA_FOLDER = '/kaggle/input/binned-dataset-v3/'\nAUX_FOLDER = '/kaggle/input/ariel-data-challenge-2024/'\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T05:04:59.808404Z","iopub.execute_input":"2025-04-20T05:04:59.80923Z","iopub.status.idle":"2025-04-20T05:04:59.814217Z","shell.execute_reply.started":"2025-04-20T05:04:59.809199Z","shell.execute_reply":"2025-04-20T05:04:59.813079Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load AIRS and FGS signal data\ndata_train = np.load(f'{DATA_FOLDER}/data_train.npy')\ndata_train_FGS = np.load(f'{DATA_FOLDER}/data_train_FGS.npy')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T05:10:09.006663Z","iopub.execute_input":"2025-04-20T05:10:09.007067Z","iopub.status.idle":"2025-04-20T05:11:13.622223Z","shell.execute_reply.started":"2025-04-20T05:10:09.007041Z","shell.execute_reply":"2025-04-20T05:11:13.620821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load target transit depth values\ntrain_solution = np.loadtxt(f'{AUX_FOLDER}/train_labels.csv', delimiter=',', skiprows=1)\ntargets = train_solution[:, 1:]\ntargets_mean = targets[:, 1:].mean(axis=1)  # average over AIRS wavelengths\nN = targets.shape[0]\n\n# Merge AIRS and FGS\nFGS_column = data_train_FGS.sum(axis=2)\ndataset = np.concatenate([data_train, FGS_column[:, :, np.newaxis, :]], axis=2)\ndataset = dataset.sum(axis=3)  # sum over y-axis to get 2D (time x wavelength)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T05:13:41.365784Z","iopub.execute_input":"2025-04-20T05:13:41.366191Z","iopub.status.idle":"2025-04-20T05:13:48.785027Z","shell.execute_reply.started":"2025-04-20T05:13:41.366158Z","shell.execute_reply":"2025-04-20T05:13:48.784121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Normalize by stellar flux (first and last 50 time steps)\ndef norm_star_spectrum(signal):\n    img_star = signal[:, :50].mean(axis=1) + signal[:, -50:].mean(axis=1)\n    return signal / img_star[:, np.newaxis, :]\n\ndataset_norm = norm_star_spectrum(dataset)\ndataset_norm = np.transpose(dataset_norm, (0, 2, 1))  # shape: (samples, wavelength, time)\n\n# Crop wavelengths to valid range\ncut_inf, cut_sup = 39, 321\nwls = np.arange(cut_sup - cut_inf + 1)\ndataset_norm = dataset_norm[:, cut_inf:cut_sup+1, :]\n\n# Split into train/valid\nN_train = 8 * N // 10\ndef split(data, N_train):\n    list_planets = random.sample(range(0, data.shape[0]), N_train)\n    list_index = np.zeros(data.shape[0], dtype=bool)\n    for i in list_planets:\n        list_index[i] = True\n    return data[list_index], data[~list_index], list_index\n\ntrain_obs, valid_obs, list_index_train = split(dataset_norm, N_train)\ntrain_targets, valid_targets = targets[list_index_train], targets[~list_index_train]\ntrain_targets_mean, valid_targets_mean = targets_mean[list_index_train], targets_mean[~list_index_train]\n\n# Normalize targets\nmin_target = train_targets_mean.min()\nmax_target = train_targets_mean.max()\ntrain_targets_mean_norm = (train_targets_mean - min_target) / (max_target - min_target)\nvalid_targets_mean_norm = (valid_targets_mean - min_target) / (max_target - min_target)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T05:14:00.548791Z","iopub.execute_input":"2025-04-20T05:14:00.549127Z","iopub.status.idle":"2025-04-20T05:14:00.899625Z","shell.execute_reply.started":"2025-04-20T05:14:00.549106Z","shell.execute_reply":"2025-04-20T05:14:00.898622Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Part 2: Model Definition - 2D CNN + LSTM hybrid model\n\ninput_shape = (dataset_norm.shape[1], dataset_norm.shape[2], 1)  # (wavelength, time, 1)\ninput_tensor = Input(shape=input_shape)\n\n# 2D CNN branch\nx_cnn = Conv2D(32, (3, 3), activation='relu', padding='same')(input_tensor)\nx_cnn = MaxPooling2D((2, 2))(x_cnn)\nx_cnn = BatchNormalization()(x_cnn)\nx_cnn = Conv2D(64, (3, 3), activation='relu', padding='same')(x_cnn)\nx_cnn = MaxPooling2D((2, 2))(x_cnn)\nx_cnn = Flatten()(x_cnn)\n\n# LSTM branch (flatten to time-major input)\nx_lstm = Reshape((dataset_norm.shape[2], dataset_norm.shape[1]))(input_tensor)  # (time, wavelength)\nx_lstm = LSTM(64, return_sequences=True)(x_lstm)\nx_lstm = LSTM(64)(x_lstm)\n\n# Merge both branches\nx = Concatenate()([x_cnn, x_lstm])\nx = Dense(128, activation='relu')(x)\nx = Dropout(0.3)(x)\nx = Dense(64, activation='relu')(x)\nx = Dropout(0.2)(x)\noutput = Dense(1, activation='linear')(x)\n\nmodel = Model(inputs=input_tensor, outputs=output)\nmodel.compile(optimizer=Adam(1e-3), loss='mse', metrics=[MeanAbsoluteError()])\nmodel.summary()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T05:14:25.907855Z","iopub.execute_input":"2025-04-20T05:14:25.908237Z","iopub.status.idle":"2025-04-20T05:14:26.434731Z","shell.execute_reply.started":"2025-04-20T05:14:25.908212Z","shell.execute_reply":"2025-04-20T05:14:26.433916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Part 3: Training and Evaluation\n\nEPOCHS = 100\nBATCH_SIZE = 16\nCHECKPOINT_PATH = 'best_model.keras'\n\ncheckpoint = ModelCheckpoint(CHECKPOINT_PATH, save_best_only=True, monitor='val_loss', mode='min')\n\nhistory = model.fit(\n    train_obs[..., np.newaxis], train_targets_mean_norm,\n    validation_data=(valid_obs[..., np.newaxis], valid_targets_mean_norm),\n    epochs=EPOCHS,\n    batch_size=BATCH_SIZE,\n    callbacks=[checkpoint],\n    verbose=1\n)\n\n# Load best model\nmodel.load_weights(CHECKPOINT_PATH)\n\n# Evaluate on validation set\nloss, mae = model.evaluate(valid_obs[..., np.newaxis], valid_targets_mean_norm)\nprint(f\"Validation Loss (MSE): {loss:.5f}\")\nprint(f\"Validation MAE: {mae:.5f}\")\n\n# Plot training history\nplt.plot(history.history['loss'], label='Training Loss')\nplt.plot(history.history['val_loss'], label='Validation Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.title('Model Training History')\nplt.grid(True)\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T05:16:41.247292Z","iopub.execute_input":"2025-04-20T05:16:41.247736Z","iopub.status.idle":"2025-04-20T06:20:04.416563Z","shell.execute_reply.started":"2025-04-20T05:16:41.247707Z","shell.execute_reply":"2025-04-20T06:20:04.415078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Part 4: Predictions and Uncertainty Estimation via MC Dropout\n\n# Function to reverse normalization\ndef unnormalize(predictions):\n    return predictions * (max_target - min_target) + min_target\n\n# Predict multiple times with dropout enabled\ndef mc_dropout_predict(model, x, n_iter=100):\n    f_model = Model(model.input, model.output)\n    preds = [f_model(x, training=True).numpy().flatten() for _ in range(n_iter)]\n    preds = np.array(preds)\n    return preds.mean(axis=0), preds.std(axis=0)\n\nmean_preds, std_preds = mc_dropout_predict(model, valid_obs[..., np.newaxis], n_iter=100)\nmean_preds = unnormalize(mean_preds)\nstd_preds = std_preds * (max_target - min_target)\n\n# True targets for comparison\ntrue_vals = valid_targets_mean\n\n# Plot prediction vs true\nplt.errorbar(np.arange(len(true_vals)), mean_preds, yerr=std_preds, fmt='o', label='Predicted')\nplt.plot(true_vals, 'r.', label='True')\nplt.xlabel('Sample Index')\nplt.ylabel('Transit Depth')\nplt.legend()\nplt.title('Predicted vs True Transit Depth (Validation Set)')\nplt.grid(True)\nplt.show()\n\n# Compute RMSE for better interpretability\nrmse = np.sqrt(np.mean((mean_preds - true_vals)**2))\nprint(f\"RMSE on validation set: {rmse*1e6:.2f} ppm\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T06:20:14.595625Z","iopub.execute_input":"2025-04-20T06:20:14.59598Z","iopub.status.idle":"2025-04-20T06:28:36.131752Z","shell.execute_reply.started":"2025-04-20T06:20:14.595944Z","shell.execute_reply":"2025-04-20T06:28:36.130771Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"NEW CODE","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom tqdm.notebook import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:03:47.935616Z","iopub.execute_input":"2025-04-22T13:03:47.936042Z","iopub.status.idle":"2025-04-22T13:03:47.940502Z","shell.execute_reply.started":"2025-04-22T13:03:47.936018Z","shell.execute_reply":"2025-04-22T13:03:47.939717Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check for GPU\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:03:51.928722Z","iopub.execute_input":"2025-04-22T13:03:51.929402Z","iopub.status.idle":"2025-04-22T13:03:52.011051Z","shell.execute_reply.started":"2025-04-22T13:03:51.929368Z","shell.execute_reply":"2025-04-22T13:03:52.010131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define constants\nWAVELENGTH_LENGTH = 55  # Number of wavelength points in the spectrum\nBATCH_SIZE = 8\nEPOCHS = 50\nLEARNING_RATE = 0.001","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:03:54.211819Z","iopub.execute_input":"2025-04-22T13:03:54.212135Z","iopub.status.idle":"2025-04-22T13:03:54.21582Z","shell.execute_reply.started":"2025-04-22T13:03:54.212114Z","shell.execute_reply":"2025-04-22T13:03:54.215068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ArielDataset(Dataset):\n    def __init__(self, data_dir, planet_ids, is_train=True):\n        self.data_dir = data_dir\n        self.is_train = is_train\n        \n        # Load wavelength grid\n        self.wavelength_grid = pd.read_csv(os.path.join(data_dir, 'wavelengths.csv'))\n        \n        # Load ADC info\n        adc_info_path = os.path.join(data_dir, 'train_adc_info.csv') if is_train else os.path.join(data_dir, 'test_adc_info.csv')\n        self.adc_info = pd.read_csv(adc_info_path)\n        \n        # Print the column names to debug\n        print(f\"ADC info columns: {self.adc_info.columns.tolist()}\")\n        \n        # Convert planet_ids to strings and filter to only include those that exist in the ADC info\n        self.adc_info['planet_id'] = self.adc_info['planet_id'].astype(str)\n        self.planet_ids = [str(pid) for pid in planet_ids if str(pid) in self.adc_info['planet_id'].values]\n        \n        if len(self.planet_ids) == 0:\n            raise ValueError(\"No valid planet IDs found in the ADC info!\")\n        \n        print(f\"Found {len(self.planet_ids)} valid planets in the dataset.\")\n        \n        # Load ground truth labels if in training mode\n        if is_train:\n            self.labels = pd.read_csv(os.path.join(data_dir, 'train_labels.csv'))\n            # Convert labels planet_id to string as well\n            self.labels['planet_id'] = self.labels['planet_id'].astype(str)\n            \n            # Print the label columns to debug\n            #print(f\"Labels columns: {self.labels.columns.tolist()}\")\n    \n    def __len__(self):\n        return len(self.planet_ids)\n    \n    def load_and_preprocess_signal(self, planet_id, instrument):\n        # Ensure planet_id is a string\n        planet_id = str(planet_id)\n        \n        # Check if planet_id exists in ADC info\n        adc_rows = self.adc_info[self.adc_info['planet_id'] == planet_id]\n        if len(adc_rows) == 0:\n            raise ValueError(f\"Planet ID {planet_id} not found in ADC info!\")\n        \n        # Get first matching row\n        planet_adc = adc_rows.iloc[0]\n        \n        # Load signal data\n        signal_path = os.path.join(self.data_dir, \"train\" if self.is_train else \"test\", planet_id, f\"{instrument}_signal.parquet\")\n        \n        # Check if file exists\n        if not os.path.exists(signal_path):\n            # Try alternative path structures\n            alt_paths = [\n                os.path.join(self.data_dir, planet_id, f\"{instrument}_signal.parquet\"),\n                os.path.join(self.data_dir, \"train\", planet_id, f\"{instrument}_signal.parquet\"),\n                os.path.join(self.data_dir, \"test\", planet_id, f\"{instrument}_signal.parquet\")\n            ]\n            \n            found_path = False\n            for path in alt_paths:\n                if os.path.exists(path):\n                    signal_path = path\n                    found_path = True\n                    break\n                    \n            if not found_path:\n                print(f\"Tried paths:\")\n                print(f\"- {os.path.join(self.data_dir, planet_id, f'{instrument}_signal.parquet')}\")\n                for path in alt_paths:\n                    print(f\"- {path}\")\n                raise FileNotFoundError(f\"Signal file not found for planet {planet_id} and instrument {instrument}\")\n            \n        # Load the data\n        print(f\"Loading signal data from: {signal_path}\")\n        signal_data = pd.read_parquet(signal_path).values\n        \n        # Get ADC conversion parameters using the correct column names\n        gain_col = f'{instrument}_adc_gain'\n        offset_col = f'{instrument}_adc_offset'\n        \n        # Check if the columns exist and use them\n        gain = 1.0  # Default gain\n        offset = 0.0  # Default offset\n        \n        if gain_col in planet_adc:\n            gain = planet_adc[gain_col]\n        else:\n            print(f\"Warning: Gain column {gain_col} not found for {instrument}. Using default gain={gain}\")\n        \n        if offset_col in planet_adc:\n            offset = planet_adc[offset_col]\n        else:\n            print(f\"Warning: Offset column {offset_col} not found for {instrument}. Using default offset={offset}\")\n        \n        # Convert to original dynamic range\n        signal_data = signal_data * gain + offset\n        \n        # Reshape data\n        if instrument == 'FGS1':\n            signal_data = signal_data.reshape(-1, 32, 32)\n        elif instrument == 'AIRS-CH0':\n            signal_data = signal_data.reshape(-1, 32, 356)\n        \n        # Find calibration directory\n        calibration_dir = os.path.join(os.path.dirname(signal_path), f\"{instrument}_calibration\")\n        if not os.path.exists(calibration_dir):\n            raise FileNotFoundError(f\"Calibration directory not found: {calibration_dir}\")\n        \n        # Load dark frame\n        dark_path = os.path.join(calibration_dir, \"dark.parquet\")\n        if not os.path.exists(dark_path):\n            raise FileNotFoundError(f\"Dark frame not found: {dark_path}\")\n        dark = pd.read_parquet(dark_path).values\n        \n        # Load flat field\n        flat_path = os.path.join(calibration_dir, \"flat.parquet\")\n        if not os.path.exists(flat_path):\n            raise FileNotFoundError(f\"Flat field not found: {flat_path}\")\n        flat = pd.read_parquet(flat_path).values\n        \n        # Reshape calibration data\n        if instrument == 'FGS1':\n            dark = dark.reshape(32, 32)\n            flat = flat.reshape(32, 32)\n        elif instrument == 'AIRS-CH0':\n            dark = dark.reshape(32, 356)\n            flat = flat.reshape(32, 356)\n        \n        # Apply basic calibration (subtract dark, divide by flat)\n        # Apply to each frame in the time series\n        calibrated_data = np.zeros_like(signal_data)\n        for i in range(signal_data.shape[0]):\n            # Subtract dark current\n            tmp = signal_data[i] - dark\n            # Apply flat field correction (avoiding division by zero)\n            calibrated_data[i] = np.divide(tmp, flat, out=np.zeros_like(tmp), where=flat!=0)\n        \n        # Add simple detrending of the time series (remove linear trend from each pixel)\n        for y in range(calibrated_data.shape[1]):\n            for x in range(calibrated_data.shape[2]):\n                pixel_ts = calibrated_data[:, y, x]\n                # Simple linear detrending\n                x_vals = np.arange(len(pixel_ts))\n                if len(pixel_ts) > 1:  # Ensure we have enough points\n                    slope, intercept = np.polyfit(x_vals, pixel_ts, 1)\n                    trend = slope * x_vals + intercept\n                    calibrated_data[:, y, x] = pixel_ts - trend\n        \n        # Convert to torch tensor and normalize\n        calibrated_tensor = torch.tensor(calibrated_data, dtype=torch.float32)\n        \n        # Normalize by subtracting mean and dividing by std\n        if torch.std(calibrated_tensor) > 0:\n            calibrated_tensor = (calibrated_tensor - torch.mean(calibrated_tensor)) / torch.std(calibrated_tensor)\n        \n        # Reshape for CNN input (add channel dimension)\n        calibrated_tensor = calibrated_tensor.unsqueeze(1)  # Shape: [time, 1, height, width]\n        \n        return calibrated_tensor\n    \n    def __getitem__(self, idx):\n        planet_id = self.planet_ids[idx]\n        \n        try:\n            # Load and preprocess FGS1 data\n            fgs1_data = self.load_and_preprocess_signal(planet_id, 'FGS1')\n            \n            # Load and preprocess AIRS-CH0 data\n            airs_data = self.load_and_preprocess_signal(planet_id, 'AIRS-CH0')\n            \n            # Subsample to manage memory (especially for FGS1 which has 135,000 frames)\n            # Take even more aggressive subsampling\n            fgs1_data = fgs1_data[::20]  # Take every 20th frame\n            airs_data = airs_data[::10]   # Take every 10th frame for AIRS-CH0\n            \n            # If training, get labels\n            if self.is_train:\n                label_rows = self.labels[self.labels['planet_id'] == planet_id]\n                if len(label_rows) == 0:\n                    raise ValueError(f\"No label found for planet ID {planet_id}\")\n                \n                planet_label = label_rows.iloc[0]\n                \n                # Print all the columns in the first row to debug\n                #print(f\"Label column names for planet {planet_id}: {planet_label.index.tolist()}\")\n                \n                # Try different approaches to extract spectrum values\n                \n                # Approach 1: Check for spectrum_0, spectrum_1, etc. columns\n                if f'spectrum_0' in planet_label:\n                    try:\n                        spectrum = np.array([planet_label[f'spectrum_{i}'] for i in range(WAVELENGTH_LENGTH)])\n                    except Exception as e:\n                        print(f\"Failed to extract spectrum with 'spectrum_i' pattern: {str(e)}\")\n                        raise\n                \n                # Approach 2: Look for a column named 'spectrum' that might contain the full array\n                elif 'spectrum' in planet_label:\n                    spectrum_data = planet_label['spectrum']\n                    if isinstance(spectrum_data, str):\n                        # Try parsing string representation of array\n                        try:\n                            spectrum = np.array(eval(spectrum_data))\n                        except:\n                            # If that fails, try other parsing methods\n                            try:\n                                spectrum = np.array([float(x) for x in spectrum_data.strip('[]').split(',')])\n                            except:\n                                raise ValueError(f\"Cannot parse spectrum string for planet {planet_id}\")\n                    else:\n                        # If it's already an array or other format\n                        spectrum = np.array(spectrum_data)\n                \n                # Approach 3: Check for wavelength-indexed columns\n                else:\n                    # Get wavelength values\n                    wavelength_cols = [col for col in self.wavelength_grid.columns if col.startswith('wl_')]\n                    wavelength_cols = wavelength_cols[:WAVELENGTH_LENGTH]\n                    \n                    if all(col in planet_label for col in wavelength_cols):\n                        spectrum = np.array([planet_label[col] for col in wavelength_cols])\n                    else:\n                        # As a last resort, just look for any columns that might be spectrum data\n                        # This is very speculative\n                        numeric_cols = [col for col in planet_label.index if col.replace('.', '').isdigit()]\n                        if len(numeric_cols) >= WAVELENGTH_LENGTH:\n                            # Sort by numeric value\n                            numeric_cols.sort(key=lambda x: float(x))\n                            spectrum = np.array([planet_label[col] for col in numeric_cols[:WAVELENGTH_LENGTH]])\n                        else:\n                            # If all else fails, just use zero placeholder\n                            print(f\"Warning: Could not find spectrum data for planet {planet_id}. Using zeros.\")\n                            spectrum = np.zeros(WAVELENGTH_LENGTH)\n                \n                spectrum_tensor = torch.tensor(spectrum, dtype=torch.float32)\n                \n                return {\n                    'planet_id': planet_id,\n                    'fgs1': fgs1_data,\n                    'airs': airs_data,\n                    'spectrum': spectrum_tensor\n                }\n            else:\n                return {\n                    'planet_id': planet_id,\n                    'fgs1': fgs1_data,\n                    'airs': airs_data\n                }\n        except Exception as e:\n            print(f\"Error processing planet {planet_id}: {str(e)}\")\n            # Instead of re-raising, return a placeholder with minimal data\n            # This allows the dataloader to continue with other samples\n            if self.is_train:\n                return {\n                    'planet_id': planet_id,\n                    'fgs1': torch.zeros((1, 1, 32, 32), dtype=torch.float32),\n                    'airs': torch.zeros((1, 1, 32, 356), dtype=torch.float32),\n                    'spectrum': torch.zeros(WAVELENGTH_LENGTH, dtype=torch.float32)\n                }\n            else:\n                return {\n                    'planet_id': planet_id,\n                    'fgs1': torch.zeros((1, 1, 32, 32), dtype=torch.float32),\n                    'airs': torch.zeros((1, 1, 32, 356), dtype=torch.float32)\n                }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:03:56.144398Z","iopub.execute_input":"2025-04-22T13:03:56.144653Z","iopub.status.idle":"2025-04-22T13:03:56.168665Z","shell.execute_reply.started":"2025-04-22T13:03:56.144634Z","shell.execute_reply":"2025-04-22T13:03:56.168039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Neural Network Architecture\nclass CNNBlock(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(CNNBlock, self).__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1)\n        self.bn = nn.BatchNorm2d(out_channels)\n        self.relu = nn.ReLU()\n        self.pool = nn.MaxPool2d(2)\n    \n    def forward(self, x):\n        x = self.conv(x)\n        x = self.bn(x)\n        x = self.relu(x)\n        x = self.pool(x)\n        return x\n\nclass CNN_LSTM_Branch(nn.Module):\n    def __init__(self, in_channels, time_steps, height, width):\n        super(CNN_LSTM_Branch, self).__init__()\n        self.time_steps = time_steps\n        \n        # CNN layers\n        self.cnn = nn.Sequential(\n            CNNBlock(in_channels, 16),\n            CNNBlock(16, 32),\n            CNNBlock(32, 64),\n        )\n        \n        # Calculate CNN output size\n        with torch.no_grad():\n            dummy_input = torch.zeros(1, in_channels, height, width)\n            dummy_output = self.cnn(dummy_input)\n            cnn_output_size = dummy_output.numel()\n        \n        # LSTM layers\n        self.lstm = nn.LSTM(\n            input_size=cnn_output_size,\n            hidden_size=128,\n            num_layers=2,\n            batch_first=True,\n            dropout=0.2\n        )\n        \n        self.output_size = 128\n    \n    def forward(self, x):\n        batch_size = x.size(0)\n        \n        # Process each timestep with CNN\n        cnn_outputs = []\n        for t in range(min(self.time_steps, x.size(1))):\n            cnn_out = self.cnn(x[:, t])\n            cnn_out = cnn_out.view(batch_size, -1)\n            cnn_outputs.append(cnn_out)\n        \n        # Stack CNN outputs along time dimension\n        cnn_sequence = torch.stack(cnn_outputs, dim=1)\n        \n        # Process with LSTM\n        lstm_out, _ = self.lstm(cnn_sequence)\n        \n        # Take output from last timestep\n        return lstm_out[:, -1, :]\n\nclass HybridFusionNet(nn.Module):\n    def __init__(self, fgs_time_steps=13500, airs_time_steps=11250):\n        super(HybridFusionNet, self).__init__()\n        \n        # Branches for each instrument\n        self.branch_fgs = CNN_LSTM_Branch(in_channels=1, time_steps=fgs_time_steps, height=32, width=32)\n        self.branch_airs = CNN_LSTM_Branch(in_channels=1, time_steps=airs_time_steps, height=32, width=356)\n        \n        # Fusion layers\n        fusion_input_size = self.branch_fgs.output_size + self.branch_airs.output_size\n        self.fusion_layer = nn.Sequential(\n            nn.Linear(fusion_input_size, 256),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, 128),\n            nn.ReLU(),\n            nn.Dropout(0.2)\n        )\n        \n        # Output layers\n        self.spectrum_layer = nn.Linear(128, WAVELENGTH_LENGTH)\n        self.uncertainty_layer = nn.Linear(128, WAVELENGTH_LENGTH)\n    \n    def forward(self, fgs, airs):\n        # Process each branch\n        fgs_features = self.branch_fgs(fgs)\n        airs_features = self.branch_airs(airs)\n        \n        # Concatenate features\n        combined_features = torch.cat((fgs_features, airs_features), dim=1)\n        \n        # Process through fusion layers\n        fused_features = self.fusion_layer(combined_features)\n        \n        # Generate outputs\n        spectrum = self.spectrum_layer(fused_features)\n        uncertainty = F.softplus(self.uncertainty_layer(fused_features))  # Ensure positive uncertainty values\n        \n        return spectrum, uncertainty\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:04:01.19109Z","iopub.execute_input":"2025-04-22T13:04:01.191673Z","iopub.status.idle":"2025-04-22T13:04:01.202288Z","shell.execute_reply.started":"2025-04-22T13:04:01.191649Z","shell.execute_reply":"2025-04-22T13:04:01.201401Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Custom loss function that takes into account uncertainty\nclass UncertaintyLoss(nn.Module):\n    def __init__(self):\n        super(UncertaintyLoss, self).__init__()\n    \n    def forward(self, pred_spectrum, true_spectrum, uncertainty):\n        # Calculate squared error\n        squared_error = (pred_spectrum - true_spectrum) ** 2\n        \n        # Weighted by inverse uncertainty (plus small value to avoid division by zero)\n        loss = torch.mean(squared_error / (uncertainty + 1e-8) + torch.log(uncertainty + 1e-8))\n        \n        return loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:04:06.228155Z","iopub.execute_input":"2025-04-22T13:04:06.228642Z","iopub.status.idle":"2025-04-22T13:04:06.232887Z","shell.execute_reply.started":"2025-04-22T13:04:06.228611Z","shell.execute_reply":"2025-04-22T13:04:06.232291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to train the model\ndef train_model(model, train_loader, val_loader, criterion, optimizer, num_epochs, device):\n    best_val_loss = float('inf')\n    train_losses = []\n    val_losses = []\n    \n    for epoch in range(num_epochs):\n        # Training phase\n        model.train()\n        running_loss = 0.0\n        \n        for batch in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} (Training)\"):\n            fgs1_data = batch['fgs1'].to(device)\n            airs_data = batch['airs'].to(device)\n            spectrum = batch['spectrum'].to(device)\n            \n            # Zero the parameter gradients\n            optimizer.zero_grad()\n            \n            # Forward pass\n            pred_spectrum, uncertainty = model(fgs1_data, airs_data)\n            \n            # Calculate loss\n            loss = criterion(pred_spectrum, spectrum, uncertainty)\n            \n            # Backward pass and optimize\n            loss.backward()\n            optimizer.step()\n            \n            running_loss += loss.item() * fgs1_data.size(0)\n        \n        epoch_train_loss = running_loss / len(train_loader.dataset)\n        train_losses.append(epoch_train_loss)\n        \n        # Validation phase\n        model.eval()\n        running_val_loss = 0.0\n        \n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{num_epochs} (Validation)\"):\n                fgs1_data = batch['fgs1'].to(device)\n                airs_data = batch['airs'].to(device)\n                spectrum = batch['spectrum'].to(device)\n                \n                # Forward pass\n                pred_spectrum, uncertainty = model(fgs1_data, airs_data)\n                \n                # Calculate loss\n                val_loss = criterion(pred_spectrum, spectrum, uncertainty)\n                \n                running_val_loss += val_loss.item() * fgs1_data.size(0)\n        \n        epoch_val_loss = running_val_loss / len(val_loader.dataset)\n        val_losses.append(epoch_val_loss)\n        \n        print(f\"Epoch {epoch+1}/{num_epochs}: Train Loss: {epoch_train_loss:.6f}, Val Loss: {epoch_val_loss:.6f}\")\n        \n        # Save best model\n        if epoch_val_loss < best_val_loss:\n            best_val_loss = epoch_val_loss\n            torch.save(model.state_dict(), 'best_model.pth')\n            print(f\"Model saved at epoch {epoch+1} with validation loss: {best_val_loss:.6f}\")\n    \n    return train_losses, val_losses","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:04:08.996391Z","iopub.execute_input":"2025-04-22T13:04:08.996634Z","iopub.status.idle":"2025-04-22T13:04:09.004335Z","shell.execute_reply.started":"2025-04-22T13:04:08.996615Z","shell.execute_reply":"2025-04-22T13:04:09.003728Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to make predictions\ndef predict(model, test_loader, device):\n    model.eval()\n    predictions = {}\n    \n    with torch.no_grad():\n        for batch in tqdm(test_loader, desc=\"Making predictions\"):\n            planet_ids = batch['planet_id']\n            fgs1_data = batch['fgs1'].to(device)\n            airs_data = batch['airs'].to(device)\n            \n            # Forward pass\n            pred_spectrum, uncertainty = model(fgs1_data, airs_data)\n            \n            # Store predictions\n            for i, planet_id in enumerate(planet_ids):\n                predictions[planet_id] = {\n                    'spectrum': pred_spectrum[i].cpu().numpy(),\n                    'uncertainty': uncertainty[i].cpu().numpy()\n                }\n    \n    return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:04:11.775367Z","iopub.execute_input":"2025-04-22T13:04:11.776077Z","iopub.status.idle":"2025-04-22T13:04:11.780988Z","shell.execute_reply.started":"2025-04-22T13:04:11.776052Z","shell.execute_reply":"2025-04-22T13:04:11.780201Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to evaluate model performance\ndef evaluate_model(model, val_loader, device):\n    model.eval()\n    all_predictions = []\n    all_targets = []\n    all_uncertainties = []\n    \n    with torch.no_grad():\n        for batch in tqdm(val_loader, desc=\"Evaluating model\"):\n            fgs1_data = batch['fgs1'].to(device)\n            airs_data = batch['airs'].to(device)\n            spectrum = batch['spectrum'].to(device)\n            \n            # Forward pass\n            pred_spectrum, uncertainty = model(fgs1_data, airs_data)\n            \n            # Store results\n            all_predictions.append(pred_spectrum.cpu().numpy())\n            all_targets.append(spectrum.cpu().numpy())\n            all_uncertainties.append(uncertainty.cpu().numpy())\n    \n    # Concatenate results\n    all_predictions = np.vstack(all_predictions)\n    all_targets = np.vstack(all_targets)\n    all_uncertainties = np.vstack(all_uncertainties)\n    \n    # Calculate metrics\n    mse = np.mean((all_predictions - all_targets) ** 2)\n    rmse = np.sqrt(mse)\n    \n    # Calculate uncertainty-weighted metrics\n    weighted_error = np.mean((all_predictions - all_targets) ** 2 / (all_uncertainties + 1e-8))\n    \n    return {\n        'MSE': mse,\n        'RMSE': rmse,\n        'Weighted_Error': weighted_error\n    }\n\n# Function to visualize results\ndef visualize_results(model, val_loader, device, wavelength_grid, num_samples=3):\n    model.eval()\n    \n    # Get random samples\n    all_results = []\n    all_planet_ids = []\n    \n    with torch.no_grad():\n        for batch in val_loader:\n            fgs1_data = batch['fgs1'].to(device)\n            airs_data = batch['airs'].to(device)\n            spectrum = batch['spectrum'].cpu().numpy()\n            planet_ids = batch['planet_id']\n            \n            # Forward pass\n            pred_spectrum, uncertainty = model(fgs1_data, airs_data)\n            pred_spectrum = pred_spectrum.cpu().numpy()\n            uncertainty = uncertainty.cpu().numpy()\n            \n            # Store results\n            for i in range(len(planet_ids)):\n                all_results.append((pred_spectrum[i], spectrum[i], uncertainty[i]))\n                all_planet_ids.append(planet_ids[i])\n            \n            if len(all_results) >= num_samples:\n                break\n    \n    # Plot results\n    wavelengths = [col for col in wavelength_grid.columns if col.startswith('wl_')]\n\n    \n    plt.figure(figsize=(15, 5 * min(num_samples, len(all_results))))\n    \n    for i in range(min(num_samples, len(all_results))):\n        pred, true, uncert = all_results[i]\n        planet_id = all_planet_ids[i]\n        \n        plt.subplot(min(num_samples, len(all_results)), 1, i+1)\n        \n        # Plot true spectrum\n        plt.plot(wavelengths, true, 'o-', label='True Spectrum', color='blue')\n        \n        # Plot predicted spectrum with uncertainty\n        plt.plot(wavelengths, pred, 'o-', label='Predicted Spectrum', color='red')\n        plt.fill_between(wavelengths, pred - uncert, pred + uncert, color='red', alpha=0.2, label='Uncertainty')\n        \n        plt.title(f'Planet ID: {planet_id}')\n        plt.xlabel('Wavelength (μm)')\n        plt.ylabel('Transit Depth')\n        plt.legend()\n        plt.grid(True, linestyle='--', alpha=0.7)\n    \n    plt.tight_layout()\n    plt.savefig('prediction_samples.png')\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:04:13.890294Z","iopub.execute_input":"2025-04-22T13:04:13.890789Z","iopub.status.idle":"2025-04-22T13:04:13.90165Z","shell.execute_reply.started":"2025-04-22T13:04:13.890767Z","shell.execute_reply":"2025-04-22T13:04:13.900997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Main execution\ndef run_pipeline():\n    # Set paths\n    data_dir = '/kaggle/input/ariel-data-challenge-2024'  # Update this to your Kaggle data directory\n    \n    # Get list of planet IDs\n    train_adc_info = pd.read_csv(os.path.join(data_dir, 'train_adc_info.csv'))\n    train_planet_ids = train_adc_info['planet_id'].unique()\n    train_planet_ids = train_planet_ids[:int(0.1 * len(train_planet_ids))]  # Use only 30%\n\n    \n    # Split train into train and validation\n    train_ids, val_ids = train_test_split(train_planet_ids, test_size=0.2, random_state=42)\n    \n    print(f\"Number of training planets: {len(train_ids)}\")\n    print(f\"Number of validation planets: {len(val_ids)}\")\n    \n    # Load wavelength grid for visualization\n    wavelength_grid = pd.read_csv(os.path.join(data_dir, 'wavelengths.csv'))\n    \n    # Create datasets\n    train_dataset = ArielDataset(data_dir, train_ids, is_train=True)\n    val_dataset = ArielDataset(data_dir, val_ids, is_train=True)\n    \n    # Create data loaders with smaller batch size for memory constraints\n    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\n    \n    # Create model - using smaller time steps to reduce memory usage\n    model = HybridFusionNet(fgs_time_steps=100, airs_time_steps=100).to(device)\n    \n    # Print model summary\n    total_params = sum(p.numel() for p in model.parameters())\n    print(f\"Total number of parameters: {total_params:,}\")\n    \n    # Create loss function and optimizer\n    criterion = UncertaintyLoss()\n    optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)\n    \n    # Train model\n    train_losses, val_losses = train_model(model, train_loader, val_loader, criterion, optimizer, EPOCHS, device)\n    \n    # Plot training history\n    plt.figure(figsize=(10, 5))\n    plt.plot(train_losses, label='Training Loss')\n    plt.plot(val_losses, label='Validation Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Training and Validation Loss')\n    plt.legend()\n    plt.savefig('training_history.png')\n    plt.show()\n    \n    # Load best model\n    model.load_state_dict(torch.load('best_model.pth'))\n    \n    # Evaluate model\n    metrics = evaluate_model(model, val_loader, device)\n    print(\"Validation Metrics:\")\n    for metric_name, metric_value in metrics.items():\n        print(f\"{metric_name}: {metric_value:.6f}\")\n    \n    # Visualize some results\n    visualize_results(model, val_loader, device, wavelength_grid, num_samples=3)\n    \n    # Make predictions on test set (if available)\n    test_adc_info_path = os.path.join(data_dir, 'test_adc_info.csv')\n    if os.path.exists(test_adc_info_path):\n        test_adc_info = pd.read_csv(test_adc_info_path)\n        test_planet_ids = test_adc_info['planet_id'].unique()\n        \n        print(f\"Number of test planets: {len(test_planet_ids)}\")\n        \n        # Create test dataset and loader\n        test_dataset = ArielDataset(data_dir, test_planet_ids, is_train=False)\n        test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\n        \n        # Make predictions\n        predictions = predict(model, test_loader, device)\n        \n        # Save predictions\n        submission = pd.DataFrame()\n        submission['planet_id'] = test_planet_ids\n        \n        for i in range(WAVELENGTH_LENGTH):\n            submission[f'spectrum_{i}'] = [predictions[pid]['spectrum'][i] for pid in test_planet_ids]\n            submission[f'uncertainty_{i}'] = [predictions[pid]['uncertainty'][i] for pid in test_planet_ids]\n        \n        submission.to_csv('submission.csv', index=False)\n        print(\"Submission file created.\")\n\ndef run_memory_efficient_pipeline():\n    \"\"\"\n    This function runs a memory-efficient version of the pipeline\n    by processing fewer planets and using more aggressive subsampling\n    \"\"\"\n    # Set paths\n    data_dir = '/kaggle/input/ariel-data-challenge-2024'  # Update this to your Kaggle data directory\n    \n    # First, explore the directory structure to understand it\n    print(f\"Data directory: {data_dir}\")\n    if os.path.exists(data_dir):\n        print(f\"Contents of data directory: {os.listdir(data_dir)}\")\n    else:\n        print(f\"Data directory does not exist: {data_dir}\")\n        return\n    \n    # Get list of planet IDs\n    train_adc_info_path = os.path.join(data_dir, 'train_adc_info.csv')\n    if not os.path.exists(train_adc_info_path):\n        print(f\"ADC info file not found: {train_adc_info_path}\")\n        return\n        \n    train_adc_info = pd.read_csv(train_adc_info_path)\n    \n    # Verify the structure of the data\n    print(\"Sample ADC info columns:\", train_adc_info.columns.tolist())\n    print(\"First few rows of ADC info:\")\n    print(train_adc_info.head())\n    \n    # Check if planet_id exists in the dataframe\n    if 'planet_id' not in train_adc_info.columns:\n        raise ValueError(\"Column 'planet_id' not found in ADC info file!\")\n    \n    # Get unique planet IDs - use only first 5 planets for memory efficiency\n    train_planet_ids = train_adc_info['planet_id'].unique()[:5]\n    print(f\"Selected planet IDs: {train_planet_ids}\")\n    \n    # Check if we need to convert to string\n    print(f\"Planet ID type: {type(train_planet_ids[0])}\")\n    \n    # Check the directory structure for the first planet\n    for planet_id in train_planet_ids[:1]:  # Just check the first one\n        planet_id_str = str(planet_id)\n        # Try different possible directory structures\n        test_paths = [\n            os.path.join(data_dir, planet_id_str),\n            os.path.join(data_dir, \"train\", planet_id_str),\n            os.path.join(data_dir, \"test\", planet_id_str)\n        ]\n        \n        for path in test_paths:\n            if os.path.exists(path):\n                print(f\"Found planet directory: {path}\")\n                print(f\"Contents: {os.listdir(path)}\")\n                # Check for signal files\n                for instrument in ['FGS1', 'AIRS-CH0']:\n                    signal_file = f\"{instrument}_signal.parquet\"\n                    if os.path.exists(os.path.join(path, signal_file)):\n                        print(f\"Found signal file: {os.path.join(path, signal_file)}\")\n                    else:\n                        print(f\"Missing signal file: {os.path.join(path, signal_file)}\")\n    \n    # Split train into train and validation\n    train_ids, val_ids = train_test_split(train_planet_ids, test_size=0.25, random_state=42)\n    \n    print(f\"Number of training planets: {len(train_ids)}\")\n    print(f\"Number of validation planets: {len(val_ids)}\")\n    \n    try:\n        # Create datasets - with much more aggressive subsampling\n        class ExtremeMemoryEfficientArielDataset(ArielDataset):\n            def __getitem__(self, idx):\n                result = super().__getitem__(idx)\n                # Keep only a tiny fraction of the data for extreme memory efficiency\n                if 'fgs1' in result:\n                    result['fgs1'] = result['fgs1'][:min(5, result['fgs1'].shape[0])]\n                if 'airs' in result:\n                    result['airs'] = result['airs'][:min(5, result['airs'].shape[0])]\n                return result\n        \n        print(\"Creating datasets with extreme memory efficiency...\")\n        train_dataset = ExtremeMemoryEfficientArielDataset(data_dir, train_ids, is_train=True)\n        val_dataset = ExtremeMemoryEfficientArielDataset(data_dir, val_ids, is_train=True)\n        \n        # Create data loaders with smaller batch size\n        train_loader = DataLoader(train_dataset, batch_size=1, shuffle=True, num_workers=0)\n        val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False, num_workers=0)\n        \n        # Create model with very small time steps\n        model = HybridFusionNet(fgs_time_steps=5, airs_time_steps=5).to(device)\n        \n        # Print model summary\n        total_params = sum(p.numel() for p in model.parameters())\n        print(f\"Total number of parameters: {total_params:,}\")\n        \n        # Create loss function and optimizer\n        criterion = UncertaintyLoss()\n        optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)\n        \n        # Train model for just 1 epoch to test\n        print(\"Training for 1 epoch as a test...\")\n        train_losses, val_losses = train_model(model, train_loader, val_loader, criterion, optimizer, 1, device)\n        \n        # Save model\n        torch.save(model.state_dict(), 'test_model.pth')\n        print(\"Model saved successfully!\")\n        \n    except Exception as e:\n        print(f\"Error in pipeline: {str(e)}\")\n        import traceback\n        traceback.print_exc()\n        \n        # Try with a single planet\n        if len(train_planet_ids) > 0:\n            try:\n                print(\"\\nTrying with just a single planet for debugging...\")\n                single_planet = [train_planet_ids[0]]\n                single_dataset = ArielDataset(data_dir, single_planet, is_train=True)\n                \n                # Try accessing one sample\n                print(\"Accessing single sample...\")\n                sample = single_dataset[0]\n                print(f\"Successfully loaded sample for planet {sample['planet_id']}\")\n                print(f\"FGS1 shape: {sample['fgs1'].shape}\")\n                print(f\"AIRS-CH0 shape: {sample['airs'].shape}\")\n                print(f\"Spectrum shape: {sample['spectrum'].shape}\")\n            except Exception as e2:\n                print(f\"Error with single planet approach: {str(e2)}\")\n                import traceback\n                traceback.print_exc()\n\n# Helper function to plot a sample spectrum\ndef plot_sample_spectrum(wavelength_grid, spectrum, title=\"Sample Spectrum\"):\n    wavelengths = wavelength_grid['wavelength'].values\n    \n    plt.figure(figsize=(10, 5))\n    plt.plot(wavelengths, spectrum, 'o-', color='blue')\n    plt.title(title)\n    plt.xlabel('Wavelength (μm)')\n    plt.ylabel('Transit Depth')\n    plt.grid(True, linestyle='--', alpha=0.7)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:04:18.122723Z","iopub.execute_input":"2025-04-22T13:04:18.123267Z","iopub.status.idle":"2025-04-22T13:04:18.143429Z","shell.execute_reply.started":"2025-04-22T13:04:18.123242Z","shell.execute_reply":"2025-04-22T13:04:18.142875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run the full pipeline or the memory efficient version based on available resources\nif __name__ == \"__main__\":\n    # Try the memory efficient version first\n    try:\n        print(\"Running memory efficient pipeline...\")\n        run_memory_efficient_pipeline()\n        \n        # If successful and memory allows, try running the full pipeline\n        print(\"\\nMemory efficient pipeline completed successfully.\")\n        print(\"Attempting to run the full pipeline...\")\n        run_pipeline()\n    except Exception as e:\n        print(f\"Error: {e}\")\n        print(\"Memory efficient pipeline failed. Trying with even more reduced parameters...\")\n        \n        # Further reduced parameters if memory is still an issue\n        # Set paths\n        data_dir = '/kaggle/input/ariel-data-challenge-2024/'\n        \n        # Get list of planet IDs - use only 5 planets\n        train_adc_info = pd.read_csv(os.path.join(data_dir, 'train_adc_info.csv'))\n        train_planet_ids = train_adc_info['planet_id'].unique()[:5]\n        \n        # Load a single planet to see if it works\n        sample_dataset = ArielDataset(data_dir, [train_planet_ids[0]], is_train=True)\n        sample = sample_dataset[0]\n        \n        # Print dimensions\n        print(f\"FGS1 data shape: {sample['fgs1'].shape}\")\n        print(f\"AIRS-CH0 data shape: {sample['airs'].shape}\")\n        \n        # Plot sample spectrum\n        wavelength_grid = pd.read_csv(os.path.join(data_dir, 'wavelengths.csv'))\n        plot_sample_spectrum(wavelength_grid, sample['spectrum'].numpy(), \"Sample Ground Truth Spectrum\")\n        \n        print(\"Data loading and visualization successful. Please adjust the pipeline parameters to fit your memory constraints.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T13:04:31.921392Z","iopub.execute_input":"2025-04-22T13:04:31.921678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom pathlib import Path\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport os\nfrom sklearn.metrics import mean_squared_error, mean_absolute_error, r2_score\n\n# Set device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\n\n# Set up paths\nDATA_DIR = Path('/kaggle/input/ariel-data-challenge-2024')\n\n# Custom Dataset class with preprocessing from host solution\nclass ArielDataset(Dataset):\n    def __init__(self, planet_ids, train=True):\n        self.planet_ids = planet_ids\n        self.train = train\n        self.base_path = DATA_DIR / ('train' if train else 'test')\n        \n        # Load metadata\n        self.adc_info = pd.read_csv(DATA_DIR / ('train_adc_info.csv' if train else 'test_adc_info.csv'))\n        if train:\n            self.labels = pd.read_csv(DATA_DIR / 'train_labels.csv')\n        \n        # Preload ADC parameters and labels\n        self.adc_params = {row['planet_id']: row for _, row in self.adc_info.iterrows()}\n        if train:\n            self.label_dict = {row['planet_id']: row[1:].values for _, row in self.labels.iterrows()}\n        \n    def __len__(self):\n        return len(self.planet_ids)\n    \n    def __getitem__(self, idx):\n        planet_id = self.planet_ids[idx]\n        \n        try:\n            # Load and process AIRS-CH0 data\n            airs_signal = pd.read_parquet(self.base_path / str(planet_id) / 'AIRS-CH0_signal.parquet')\n            airs_signal = torch.FloatTensor(airs_signal.values.astype(np.float32) / 65535.0)\n            \n            # Load and process FGS1 data\n            fgs1_signal = pd.read_parquet(self.base_path / str(planet_id) / 'FGS1_signal.parquet')\n            fgs1_signal = torch.FloatTensor(fgs1_signal.values.astype(np.float32) / 65535.0)\n            \n            # Get ADC parameters\n            adc_params = self.adc_params[planet_id]\n            \n            # Apply ADC correction\n            airs_signal = airs_signal * adc_params['AIRS-CH0_adc_gain'] + adc_params['AIRS-CH0_adc_offset']\n            fgs1_signal = fgs1_signal * adc_params['FGS1_adc_gain'] + adc_params['FGS1_adc_offset']\n            \n            # Reshape signals\n            airs_signal = airs_signal.reshape(-1, 32, 356)\n            fgs1_signal = fgs1_signal.reshape(-1, 32, 32)\n            \n            # Normalize using star spectrum (first and last 50 instants)\n            img_star = airs_signal[:,:50].mean(dim=1) + airs_signal[:,-50:].mean(dim=1)\n            airs_signal = airs_signal / img_star.unsqueeze(1)\n            \n            # Cut transit (ingress=75, egress=115)\n            airs_signal = airs_signal[:,75:115]\n            \n            # Remove mean value\n            airs_signal = airs_signal - airs_signal.mean(dim=1, keepdim=True)\n            \n            if self.train:\n                labels = torch.FloatTensor(self.label_dict[planet_id])\n                return airs_signal, fgs1_signal, labels\n            else:\n                return airs_signal, fgs1_signal\n                \n        except Exception as e:\n            print(f\"Error processing planet {planet_id}: {str(e)}\")\n            if self.train:\n                return torch.zeros((1, 40, 356)), torch.zeros((1, 32, 32)), torch.zeros(283)\n            else:\n                return torch.zeros((1, 40, 356)), torch.zeros((1, 32, 32))\n\n# 1D-CNN for Transit Depth (from host solution)\nclass TransitDepthNet(nn.Module):\n    def __init__(self):\n        super(TransitDepthNet, self).__init__()\n        \n        self.conv1 = nn.Sequential(\n            nn.Conv1d(1, 32, kernel_size=3, padding=1),\n            nn.BatchNorm1d(32),\n            nn.ReLU(),\n            nn.MaxPool1d(2)\n        )\n        \n        self.conv2 = nn.Sequential(\n            nn.Conv1d(32, 64, kernel_size=3, padding=1),\n            nn.BatchNorm1d(64),\n            nn.ReLU(),\n            nn.MaxPool1d(2)\n        )\n        \n        self.conv3 = nn.Sequential(\n            nn.Conv1d(64, 128, kernel_size=3, padding=1),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.MaxPool1d(2)\n        )\n        \n        self.conv4 = nn.Sequential(\n            nn.Conv1d(128, 256, kernel_size=3, padding=1),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.MaxPool1d(2)\n        )\n        \n        # Calculate the size after convolutions\n        # Input shape: (batch_size, 1, 40) -> after 4 maxpools: (batch_size, 256, 2)\n        self.fc = nn.Sequential(\n            nn.Linear(256 * 2, 500),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(500, 100),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            nn.Linear(100, 1)\n        )\n        \n    def forward(self, x):\n        # Ensure correct input shape: [batch, channels, sequence_length]\n        if len(x.shape) == 2:  # If [batch, sequence]\n            x = x.unsqueeze(1)  # Add channel dimension: [batch, 1, sequence]\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.conv3(x)\n        x = self.conv4(x)\n        x = x.view(x.size(0), -1)\n        x = self.fc(x)\n        return x\n\n# 2D-CNN for Atmospheric Features (from host solution)\nclass AtmosphericFeaturesNet(nn.Module):\n    def __init__(self):\n        super(AtmosphericFeaturesNet, self).__init__()\n        \n        self.conv1 = nn.Sequential(\n            nn.Conv2d(1, 32, kernel_size=(3,1), padding=(1,0)),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            nn.MaxPool2d((2,1))\n        )\n        \n        self.conv2 = nn.Sequential(\n            nn.Conv2d(32, 64, kernel_size=(3,1), padding=(1,0)),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.MaxPool2d((2,1))\n        )\n        \n        self.conv3 = nn.Sequential(\n            nn.Conv2d(64, 128, kernel_size=(3,1), padding=(1,0)),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.MaxPool2d((2,1))\n        )\n        \n        self.conv4 = nn.Sequential(\n            nn.Conv2d(128, 256, kernel_size=(3,1), padding=(1,0)),\n            nn.BatchNorm2d(256),\n            nn.ReLU()\n        )\n        \n        self.conv5 = nn.Sequential(\n            nn.Conv2d(256, 32, kernel_size=(1,3), padding=(0,1)),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            nn.MaxPool2d((1,2))\n        )\n        \n        self.conv6 = nn.Sequential(\n            nn.Conv2d(32, 64, kernel_size=(1,3), padding=(0,1)),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.MaxPool2d((1,2))\n        )\n        \n        self.conv7 = nn.Sequential(\n            nn.Conv2d(64, 128, kernel_size=(1,3), padding=(0,1)),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.MaxPool2d((1,2))\n        )\n        \n        self.conv8 = nn.Sequential(\n            nn.Conv2d(128, 256, kernel_size=(1,3), padding=(0,1)),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.MaxPool2d((1,2))\n        )\n        \n        # Adjust the input size for the linear layer based on actual dimensions\n        self.fc = nn.Sequential(\n            nn.Linear(256 * 2 * 22, 700),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(700, 283)\n        )\n        \n    def forward(self, x):\n        # Ensure correct input shape: [batch, channels, height, width]\n        if len(x.shape) == 3:  # If [batch, height, width]\n            x = x.unsqueeze(1)  # Add channel dimension: [batch, 1, height, width]\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.conv3(x)\n        x = self.conv4(x)\n        x = self.conv5(x)\n        x = self.conv6(x)\n        x = self.conv7(x)\n        x = self.conv8(x)\n        x = x.view(x.size(0), -1)\n        x = self.fc(x)\n        return x\n\n# Combined Model\nclass CombinedSpectraNet(nn.Module):\n    def __init__(self):\n        super(CombinedSpectraNet, self).__init__()\n        self.transit_net = TransitDepthNet()\n        self.atmospheric_net = AtmosphericFeaturesNet()\n        \n    def forward(self, time_series, spatial_data):\n        print(f\"time_series shape before squeeze: {time_series.shape}\")\n    \n        # Remove zero-sized dimensions\n        time_series = time_series.squeeze()\n    \n        print(f\"time_series shape after squeeze: {time_series.shape}\")\n    \n        # Now unpack safely\n        if len(time_series.shape) == 3:\n            batch_size, num_samples, num_features = time_series.shape\n    \n            # Reshape to [batch, 1, sequence_length]\n            white_light_curve = time_series.reshape(batch_size, 1, -1)\n    \n            # Now pass it through the conv layers\n            transit_output = self.transit_net(white_light_curve)\n        else:\n            raise ValueError(f\"Unexpected shape after squeeze: {time_series.shape}\")\n\n\n        \n        # Important fix: Ensure the shape is [batch, 1, sequence_length] for Conv1D\n        # No need to add an extra dimension - it will be added in TransitDepthNet.forward()\n        \n        # Get predictions\n        transit_output = self.transit_net(white_light_curve)  # shape: (batch_size, 1)\n        atmospheric_output = self.atmospheric_net(spatial_data)  # shape: (batch_size, 283)\n        \n        # Expand transit output to match atmospheric output\n        transit_output = transit_output.expand(-1, 283)  # shape: (batch_size, 283)\n        \n        # Combine outputs\n        final_spectra = atmospheric_output + transit_output\n        return final_spectra, transit_output, atmospheric_output\n\n# Function to plot training history\ndef plot_training_history(train_losses, val_losses, train_metrics, val_metrics):\n    fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n    \n    # Plot losses\n    axes[0, 0].plot(train_losses, label='Train Loss')\n    axes[0, 0].plot(val_losses, label='Validation Loss')\n    axes[0, 0].set_title('Training and Validation Loss')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].legend()\n    axes[0, 0].grid(True)\n    \n    # Plot MSE\n    axes[0, 1].plot(train_metrics['mse'], label='Train MSE')\n    axes[0, 1].plot(val_metrics['mse'], label='Validation MSE')\n    axes[0, 1].set_title('MSE')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('MSE')\n    axes[0, 1].legend()\n    axes[0, 1].grid(True)\n    \n    # Plot MAE\n    axes[1, 0].plot(train_metrics['mae'], label='Train MAE')\n    axes[1, 0].plot(val_metrics['mae'], label='Validation MAE')\n    axes[1, 0].set_title('MAE')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('MAE')\n    axes[1, 0].legend()\n    axes[1, 0].grid(True)\n    \n    # Plot R2\n    axes[1, 1].plot(train_metrics['r2'], label='Train R2')\n    axes[1, 1].plot(val_metrics['r2'], label='Validation R2')\n    axes[1, 1].set_title('R2 Score')\n    axes[1, 1].set_xlabel('Epoch')\n    axes[1, 1].set_ylabel('R2')\n    axes[1, 1].legend()\n    axes[1, 1].grid(True)\n    \n    plt.tight_layout()\n    plt.savefig('training_history.png')\n    plt.show()\n\n# Function to visualize predictions\ndef visualize_predictions(model, data_loader, wavelength_values):\n    model.eval()\n    \n    # Get a batch of data\n    for airs_signal, fgs1_signal, labels in data_loader:\n        airs_signal = airs_signal.to(device)\n        labels = labels.to(device)\n        \n        # Get predictions\n        with torch.no_grad():\n            final_spectra, transit_depth, atmospheric_features = model(airs_signal, airs_signal)\n        \n        # Convert to numpy arrays\n        predictions = final_spectra.cpu().numpy()\n        targets = labels.cpu().numpy()\n        \n        # Plot for the first sample in the batch\n        plt.figure(figsize=(10, 6))\n        plt.plot(wavelength_values, targets[0], label='Ground Truth')\n        plt.plot(wavelength_values, predictions[0], label='Prediction')\n        plt.xlabel('Wavelength (μm)')\n        plt.ylabel('Transit Depth')\n        plt.title('Transit Spectrum Prediction vs Ground Truth')\n        plt.legend()\n        plt.grid(True)\n        plt.savefig('prediction_visualization.png')\n        plt.show()\n        \n        # Plot the difference\n        plt.figure(figsize=(10, 6))\n        plt.plot(wavelength_values, predictions[0] - targets[0])\n        plt.xlabel('Wavelength (μm)')\n        plt.ylabel('Prediction - Ground Truth')\n        plt.title('Prediction Error')\n        plt.grid(True)\n        plt.savefig('prediction_error.png')\n        plt.show()\n        \n        # Only visualize the first sample\n        break\n\n# Training function with MC Dropout\ndef train_models(model, train_loader, val_loader, num_epochs=20, num_dropout=5):\n    criterion = nn.MSELoss()\n    optimizer = optim.Adam(model.parameters(), lr=0.001)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=3)\n    \n    train_losses = []\n    val_losses = []\n    train_metrics = {'mse': [], 'mae': [], 'r2': []}\n    val_metrics = {'mse': [], 'mae': [], 'r2': []}\n    \n    for epoch in range(num_epochs):\n        # Training phase\n        model.train()\n        train_loss = 0\n        train_preds = []\n        train_targets = []\n        \n        for airs_signal, fgs1_signal, labels in tqdm(train_loader):\n            # airs_signal shape: (batch_size, 40, 356)\n            # fgs1_signal shape: (batch_size, 32, 32)\n            # labels shape: (batch_size, 283)\n            \n            airs_signal = airs_signal.to(device)\n            labels = labels.to(device)\n            \n            optimizer.zero_grad()\n            \n            # Debug: Print shapes \n            if epoch == 0 and train_loss == 0:\n                print(f\"Input airs_signal shape: {airs_signal.shape}\")\n                white_light_curve = airs_signal.mean(dim=2)\n                print(f\"White light curve shape: {white_light_curve.shape}\")\n            \n            final_spectra, _, _ = model(airs_signal, airs_signal)\n            \n            loss = criterion(final_spectra, labels)\n            loss.backward()\n            optimizer.step()\n            \n            train_loss += loss.item()\n            train_preds.extend(final_spectra.cpu().detach().numpy())\n            train_targets.extend(labels.cpu().numpy())\n        \n        train_loss /= len(train_loader)\n        train_losses.append(train_loss)\n        \n        # Calculate training metrics\n        train_preds = np.array(train_preds)\n        train_targets = np.array(train_targets)\n        train_metrics['mse'].append(mean_squared_error(train_targets, train_preds))\n        train_metrics['mae'].append(mean_absolute_error(train_targets, train_preds))\n        train_metrics['r2'].append(r2_score(train_targets, train_preds))\n        \n        # Validation phase with MC Dropout\n        model.eval()\n        val_loss = 0\n        val_preds = []\n        val_targets = []\n        \n        with torch.no_grad():\n            for airs_signal, fgs1_signal, labels in val_loader:\n                airs_signal = airs_signal.to(device)\n                labels = labels.to(device)\n                \n                # MC Dropout\n                predictions = []\n                for _ in range(num_dropout):\n                    final_spectra, _, _ = model(airs_signal, airs_signal)\n                    predictions.append(final_spectra.cpu().numpy())\n                \n                predictions = np.array(predictions)\n                mean_pred = np.mean(predictions, axis=0)\n                std_pred = np.std(predictions, axis=0)\n                \n                loss = criterion(torch.FloatTensor(mean_pred).to(device), labels)\n                val_loss += loss.item()\n                val_preds.extend(mean_pred)\n                val_targets.extend(labels.cpu().numpy())\n        \n        val_loss /= len(val_loader)\n        val_losses.append(val_loss)\n        \n        # Calculate validation metrics\n        val_preds = np.array(val_preds)\n        val_targets = np.array(val_targets)\n        val_metrics['mse'].append(mean_squared_error(val_targets, val_preds))\n        val_metrics['mae'].append(mean_absolute_error(val_targets, val_preds))\n        val_metrics['r2'].append(r2_score(val_targets, val_preds))\n        \n        print(f'Epoch {epoch+1}/{num_epochs}:')\n        print(f'Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}')\n        print(f'Train Metrics - MSE: {train_metrics[\"mse\"][-1]:.4f}, MAE: {train_metrics[\"mae\"][-1]:.4f}, R2: {train_metrics[\"r2\"][-1]:.4f}')\n        print(f'Val Metrics - MSE: {val_metrics[\"mse\"][-1]:.4f}, MAE: {val_metrics[\"mae\"][-1]:.4f}, R2: {val_metrics[\"r2\"][-1]:.4f}')\n        \n        scheduler.step(val_loss)\n    \n    return train_losses, val_losses, train_metrics, val_metrics\n\n# Main execution\ndef main():\n    torch.cuda.empty_cache()\n    \n    # Load data\n    train_adc_info = pd.read_csv(DATA_DIR / 'train_adc_info.csv')\n    planet_ids = train_adc_info['planet_id'].values\n    \n    # Use a smaller subset for testing\n    planet_ids = planet_ids[:100]  # Limit to first 100 planets\n    \n    train_ids, val_ids = train_test_split(planet_ids, test_size=0.2, random_state=42)\n    \n    # Create datasets\n    train_dataset = ArielDataset(train_ids, train=True)\n    val_dataset = ArielDataset(val_ids, train=True)\n    \n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=2,\n        shuffle=True,\n        num_workers=0,\n        pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=2,\n        shuffle=False,\n        num_workers=0,\n        pin_memory=True\n    )\n    \n    # Get sample batch to check shapes\n    sample_airs, sample_fgs, sample_labels = next(iter(train_loader))\n    print(f\"Sample airs shape: {sample_airs.shape}\")\n    print(f\"Sample fgs shape: {sample_fgs.shape}\")\n    print(f\"Sample labels shape: {sample_labels.shape}\")\n    \n    # Initialize model\n    model = CombinedSpectraNet().to(device)\n    \n    # Train model\n    print(\"Starting training...\")\n    train_losses, val_losses, train_metrics, val_metrics = train_models(\n        model,\n        train_loader,\n        val_loader,\n        num_epochs=5,\n        num_dropout=5\n    )\n    \n    # Load wavelength values\n    wavelength = pd.read_csv(DATA_DIR / 'wavelengths.csv')\n    wavelength_values = wavelength.iloc[0].values\n    \n    # Plot training history\n    print(\"\\nPlotting training history...\")\n    plot_training_history(train_losses, val_losses, train_metrics, val_metrics)\n    \n    # Visualize predictions\n    print(\"\\nVisualizing predictions...\")\n    visualize_predictions(model, val_loader, wavelength_values)\n    \n    # Print final metrics\n    print(\"\\nFinal Metrics:\")\n    print(f\"Training - MSE: {train_metrics['mse'][-1]:.6f}, MAE: {train_metrics['mae'][-1]:.6f}, R2: {train_metrics['r2'][-1]:.4f}\")\n    print(f\"Validation - MSE: {val_metrics['mse'][-1]:.6f}, MAE: {val_metrics['mae'][-1]:.6f}, R2: {val_metrics['r2'][-1]:.4f}\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T14:54:54.55532Z","iopub.execute_input":"2025-04-21T14:54:54.556981Z","iopub.status.idle":"2025-04-21T14:55:46.149515Z","shell.execute_reply.started":"2025-04-21T14:54:54.556934Z","shell.execute_reply":"2025-04-21T14:55:46.147709Z"}},"outputs":[],"execution_count":null}]}