{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":70367,"databundleVersionId":9188054,"sourceType":"competition"},{"sourceId":191155259,"sourceType":"kernelVersion"},{"sourceId":192682085,"sourceType":"kernelVersion"}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install pytorchvideo","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:39:20.265666Z","iopub.execute_input":"2024-08-20T04:39:20.266634Z","iopub.status.idle":"2024-08-20T04:39:42.349391Z","shell.execute_reply.started":"2024-08-20T04:39:20.266601Z","shell.execute_reply":"2024-08-20T04:39:42.348326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset,DataLoader\nfrom tqdm import tqdm\nimport polars as pl\nimport os\nimport pickle\nimport matplotlib.pyplot as plt\nimport matplotlib.animation as animation\nimport torchaudio\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport random\nimport time\nimport math\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:39:42.351243Z","iopub.execute_input":"2024-08-20T04:39:42.351566Z","iopub.status.idle":"2024-08-20T04:39:46.5518Z","shell.execute_reply.started":"2024-08-20T04:39:42.351538Z","shell.execute_reply":"2024-08-20T04:39:46.550823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# config","metadata":{}},{"cell_type":"code","source":"class config:\n    # Paths\n    BASE_PATH = '/kaggle/input/ariel-data-challenge-2024/'\n    TRAIN_PATH = BASE_PATH + 'train/'\n    TEST_PATH = BASE_PATH + 'test/'\n    AXIS_INFO_PATH = BASE_PATH + 'axis_info.parquet'\n    SAMPLE_SUBMISSION_PATH = BASE_PATH + 'sample_submission.csv'\n    TRAIN_ADC_INFO_PATH = BASE_PATH + 'train_adc_info.csv'\n    TEST_ADC_INFO_PATH = BASE_PATH + 'test_adc_info.csv'\n    TRAIN_LABELS_PATH = BASE_PATH + 'train_labels.csv'\n    WAVELENGTHS_PATH = BASE_PATH + 'wavelengths.csv'\n    CALIBRATION_FILES = {\n        'FGS1': {\n            'dark': BASE_PATH + 'train/{planet_id}/FGS1_calibration/dark.parquet',\n            'dead': BASE_PATH + 'train/{planet_id}/FGS1_calibration/dead.parquet',\n            'flat': BASE_PATH + 'train/{planet_id}/FGS1_calibration/flat.parquet',\n            'linear_corr': BASE_PATH + 'train/{planet_id}/FGS1_calibration/linear_corr.parquet',\n            'read': BASE_PATH + 'train/{planet_id}/FGS1_calibration/read.parquet',\n        },\n        'AIRS-CH0': {\n            'dark': BASE_PATH + 'train/{planet_id}/AIRS-CH0_calibration/dark.parquet',\n            'dead': BASE_PATH + 'train/{planet_id}/AIRS-CH0_calibration/dead.parquet',\n            'flat': BASE_PATH + 'train/{planet_id}/AIRS-CH0_calibration/flat.parquet',\n            'linear_corr': BASE_PATH + 'train/{planet_id}/AIRS-CH0_calibration/linear_corr.parquet',\n            'read': BASE_PATH + 'train/{planet_id}/AIRS-CH0_calibration/read.parquet',\n        }\n    }\n    GRADIENT_ACCUMULATION_STEPS = 1\n    MAX_GRAD_NORM = 1e7\n    \n    EPOCHS = 5\n    mode = 'mean' # SNR or mean\n    AMP = True\n    OUTPUT_DIR = '/kaggle/working/v1'\nbatch_size = 16\ninstrument = None  # 'AIRS-CH0' or 'FGS1' or None\nos.makedirs(config.OUTPUT_DIR, exist_ok=True)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-08-20T04:41:03.344682Z","iopub.execute_input":"2024-08-20T04:41:03.345282Z","iopub.status.idle":"2024-08-20T04:41:03.353795Z","shell.execute_reply.started":"2024-08-20T04:41:03.345249Z","shell.execute_reply":"2024-08-20T04:41:03.352932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s: float):\n    \"Convert to minutes.\"\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since: float, percent: float):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\n\ndef get_logger(filename=config.OUTPUT_DIR+'/log'):\n    from logging import getLogger, INFO, StreamHandler, FileHandler, Formatter\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=f\"{filename}.log\")\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:04.934978Z","iopub.execute_input":"2024-08-20T04:41:04.93565Z","iopub.status.idle":"2024-08-20T04:41:04.945456Z","shell.execute_reply.started":"2024-08-20T04:41:04.935617Z","shell.execute_reply":"2024-08-20T04:41:04.94459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LOGGER = get_logger()\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:05.549024Z","iopub.execute_input":"2024-08-20T04:41:05.549845Z","iopub.status.idle":"2024-08-20T04:41:05.554339Z","shell.execute_reply.started":"2024-08-20T04:41:05.549796Z","shell.execute_reply":"2024-08-20T04:41:05.553357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_snr(signal, noise):\n    # Assuming signal and noise are 2D arrays: (time_steps, height, width)\n    signal_mean = signal.mean(axis=(1, 2))\n    noise_std = noise.std(axis=(0, 1))\n    snr = signal_mean / noise_std\n    return snr\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:06.625677Z","iopub.execute_input":"2024-08-20T04:41:06.626559Z","iopub.status.idle":"2024-08-20T04:41:06.631258Z","shell.execute_reply.started":"2024-08-20T04:41:06.626528Z","shell.execute_reply":"2024-08-20T04:41:06.630265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.mode == 'SNR':\n    def load_and_correct_signal(file_path, adc_info, calibration_files, instrument, planet_id):\n        # Load signal data\n        signal_data = pd.read_parquet(file_path).values\n\n        # Load calibration data\n        calib_files = calibration_files[instrument]\n        dark_frame = pd.read_parquet(calib_files['dark'].format(planet_id=planet_id)).values\n        flat_frame = pd.read_parquet(calib_files['flat'].format(planet_id=planet_id)).values\n        read_noise = pd.read_parquet(calib_files['read'].format(planet_id=planet_id)).values\n\n        # Load ADC info\n        adc_row = adc_info[adc_info['planet_id'] == int(planet_id)].iloc[0]\n        gain = adc_row[f'{instrument}_adc_gain']\n        offset = adc_row[f'{instrument}_adc_offset']\n\n        # Restore dynamic range\n        corrected_signal = (signal_data * gain) + offset\n\n        if instrument == \"FGS1\":\n            # Reshape to (135000, 32, 32)\n            corrected_signal = corrected_signal.reshape(135000, 32, 32)\n            # Compute SNR\n            noise = dark_frame\n            snr = compute_snr(corrected_signal, noise)\n            # Select top 90 time steps based on SNR\n            top_indices = np.argsort(snr)[-90:]\n            corrected_signal = corrected_signal[top_indices]\n\n        elif instrument == \"AIRS-CH0\":\n            # Reshape to (11250, 32, 356)\n            corrected_signal = corrected_signal.reshape(11250, 32, 356)\n            # Compute SNR\n            noise = dark_frame\n            snr = compute_snr(corrected_signal, noise)\n            # Select top 90 time steps based on SNR\n            top_indices = np.argsort(snr)[-90:]\n            corrected_signal = corrected_signal[top_indices]\n\n        else:\n            raise ValueError(\"Invalid instrument specified\")\n\n        # Apply dark, flat, and read noise corrections\n        corrected_signal = corrected_signal - dark_frame\n        corrected_signal = corrected_signal / (flat_frame - read_noise)\n\n        return corrected_signal\nif config.mode == 'mean':\n    def load_and_correct_signal(file_path, adc_info, calibration_files, instrument, planet_id):\n        # Load signal data\n        signal_data = pd.read_parquet(file_path).values\n\n        # Load calibration data\n        calib_files = calibration_files[instrument]\n        dark_frame = pd.read_parquet(calib_files['dark'].format(planet_id = planet_id)).values\n        flat_frame = pd.read_parquet(calib_files['flat'].format(planet_id = planet_id)).values\n        read_noise = pd.read_parquet(calib_files['read'].format(planet_id = planet_id)).values\n        # Apply calibration (you might need to adjust depending on calibration data formats)\n\n        # Load ADC info\n        adc_row = adc_info[adc_info['planet_id'] == int(planet_id)].iloc[0]\n        gain = adc_row[f'{instrument}_adc_gain']\n        offset = adc_row[f'{instrument}_adc_offset']\n\n        # Restore dynamic range\n        corrected_signal = (signal_data * gain) + offset\n\n        if instrument == \"FGS1\":\n            # Reshape to (135000, 32, 32)\n            corrected_signal = corrected_signal.reshape(135000, 32, 32)\n            # Average pooling to reduce time steps to 90\n            corrected_signal = F.avg_pool3d(\n                torch.tensor(corrected_signal).unsqueeze(0).unsqueeze(0).float(), \n                kernel_size=(135000 // 90, 1, 1),  # Reduce time steps to 90\n                stride=(135000 // 90,1,1)\n            ).squeeze(0).squeeze(0).numpy()\n            # The resulting shape should be (90, 32, 32)\n\n        elif instrument == \"AIRS-CH0\":\n            # Reshape to (11250, 32, 356)\n            corrected_signal = corrected_signal.reshape(11250, 32, 356)\n            # Average pooling to reduce time steps to 90\n            corrected_signal = F.avg_pool3d(\n                torch.tensor(corrected_signal).unsqueeze(0).unsqueeze(0).float(), \n                kernel_size=(11250 // 90, 1, 1),  # Reduce time steps to 90\n                stride=(11250 // 90,1,1)\n            ).squeeze(0).squeeze(0).numpy()\n        else:\n            raise ValueError(\"Invalid instrument specified\")\n\n        # Apply dark, flat, and read noise corrections\n        # Assuming dark_frame, flat_frame, and read_noise are properly shaped and compatible\n        corrected_signal = corrected_signal - dark_frame\n        corrected_signal = corrected_signal / (flat_frame - read_noise)  # Assuming flat_frame is a calibration factor\n\n        return corrected_signal\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:07.345551Z","iopub.execute_input":"2024-08-20T04:41:07.345951Z","iopub.status.idle":"2024-08-20T04:41:07.368689Z","shell.execute_reply.started":"2024-08-20T04:41:07.345904Z","shell.execute_reply":"2024-08-20T04:41:07.367719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImageDomainDataset(Dataset):\n    def __init__(self, planet_ids, dataset, instrument, transform=None):\n        \"\"\"\n        Args:\n            planet_ids (list): List of planet IDs to load data for.\n            dataset (str): 'train' or 'test'\n            instrument (str): 'AIRS-CH0' or 'FGS1'\n            transform (callable, optional): Optional transform to be applied on a sample\n        \"\"\"\n        self.dataset = dataset\n        self.instrument = instrument\n        self.transform = transform\n        self.planet_ids = planet_ids\n        # Load ADC information\n        adc_info_path = config.TRAIN_ADC_INFO_PATH if self.dataset == 'train' else config.TEST_ADC_INFO_PATH\n        self.adc_info = pd.read_csv(adc_info_path)\n        \n        # Load labels\n        if self.dataset == 'train':\n            labels_path = config.TRAIN_LABELS_PATH\n            self.labels = pd.read_csv(labels_path)\n    def normalize(self, signal):\n        \"\"\"\n        Normalize the signal data.\n        \"\"\"\n        mean = signal.mean()\n        std = signal.std()\n        return (signal - mean) / std\n    def load_signal_data(self, planet_id):\n        \"\"\"\n        Load and preprocess signal data for a specific planet ID.\n        \"\"\"\n        if self.instrument is None:\n            signals = []\n            for instr in ['AIRS-CH0', 'FGS1']:\n                file_path = f'{config.TRAIN_PATH if self.dataset == \"train\" else config.TEST_PATH}/{planet_id}/{instr}_signal.parquet'\n                signal = torch.tensor(load_and_correct_signal(file_path, self.adc_info, config.CALIBRATION_FILES, instr, planet_id), dtype=torch.float32)\n                signals.append(signal)\n\n            # Concatenate signals along the last dimension\n            signals = torch.cat(signals, dim=-1)  # Concatenating along the last dimension\n            signal = self.normalize(signal)\n            return signals\n        else:\n            file_path = f'{config.TRAIN_PATH if self.dataset == \"train\" else config.TEST_PATH}/{planet_id}/{self.instrument}_signal.parquet'\n            signal = load_and_correct_signal(file_path, self.adc_info, config.CALIBRATION_FILES, self.instrument, planet_id)\n            signal = self.normalize(signal)\n            return torch.tensor(signal, dtype=torch.float32)\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        sample = self.load_signal_data(planet_id)\n\n        # Get label for the given planet_id\n\n        if self.transform:\n            sample = self.transform(sample)\n        if self.dataset == 'test':\n            return {'signal': sample}    \n        \n        label = self.labels[self.labels.planet_id==int(planet_id)].iloc[:,1:].values.astype(float)\n\n        return {'signal': sample, 'label': torch.tensor(label, dtype=torch.float32)}","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:07.474925Z","iopub.execute_input":"2024-08-20T04:41:07.475511Z","iopub.status.idle":"2024-08-20T04:41:07.488504Z","shell.execute_reply.started":"2024-08-20T04:41:07.47548Z","shell.execute_reply":"2024-08-20T04:41:07.487596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold\ndef get_kfold_splits(planet_ids, n_splits=5):\n    \"\"\"\n    Generate k-fold splits for planet IDs.\n    \"\"\"\n    kf = KFold(n_splits=n_splits, shuffle=True, random_state=42)\n    return list(kf.split(planet_ids))\n\nall_planet_ids = [d for d in os.listdir(config.TRAIN_PATH) if os.path.isdir(os.path.join(config.TRAIN_PATH, d))]\nsplits = get_kfold_splits(all_planet_ids)\n\nfor fold, (train_indices, val_indices) in enumerate(splits):\n    train_planet_ids = [all_planet_ids[i] for i in train_indices]\n    val_planet_ids = [all_planet_ids[i] for i in val_indices]","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:07.634502Z","iopub.execute_input":"2024-08-20T04:41:07.634748Z","iopub.status.idle":"2024-08-20T04:41:08.302729Z","shell.execute_reply.started":"2024-08-20T04:41:07.634727Z","shell.execute_reply":"2024-08-20T04:41:08.301916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_dataloaders_for_fold(train_ids, valid_ids, instrument,):\n    train_dataset = ImageDomainDataset(dataset = 'train' ,planet_ids=train_ids, instrument=instrument)\n    valid_dataset = ImageDomainDataset(dataset = 'train' ,planet_ids=valid_ids, instrument=instrument)\n\n    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=4)\n    valid_loader = DataLoader(valid_dataset, batch_size=batch_size, shuffle=False, num_workers=4)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:08.304372Z","iopub.execute_input":"2024-08-20T04:41:08.304699Z","iopub.status.idle":"2024-08-20T04:41:08.310651Z","shell.execute_reply.started":"2024-08-20T04:41:08.304673Z","shell.execute_reply":"2024-08-20T04:41:08.309714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Generate Video","metadata":{}},{"cell_type":"code","source":"def save_animation(frames, filename='animation.mp4', interval=50):\n    \"\"\"\n    Save the animation of the frames to a file.\n    \n    Args:\n        frames (numpy.ndarray): Array of frames with shape (num_frames, height, width).\n        filename (str): Name of the file to save the animation.\n        interval (int): Time between frames in milliseconds.\n    \"\"\"\n    fig, ax = plt.subplots(figsize=(10, 5))\n    ax.set_title('Exoplanet Signal Animation')\n    \n    ims = []\n    for i in tqdm(range(frames.shape[0])):\n        im = ax.imshow(frames[i], cmap='gray', animated=True)\n        ims.append([im])\n    \n    ani = animation.ArtistAnimation(fig, ims, interval=interval, blit=True, repeat=True)\n    ani.save(filename, writer='ffmpeg', fps=20)\n    plt.close(fig)\n    print(f\"Animation saved as {filename}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:08.311672Z","iopub.execute_input":"2024-08-20T04:41:08.311949Z","iopub.status.idle":"2024-08-20T04:41:08.321062Z","shell.execute_reply.started":"2024-08-20T04:41:08.311927Z","shell.execute_reply":"2024-08-20T04:41:08.320259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sample = image_dataset[1].numpy()\n\n# frames = sample[::100]\n\n# # Create and show the animation\n# save_animation(frames, filename='exoplanet_signal.mp4', interval=100)","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:08.322708Z","iopub.execute_input":"2024-08-20T04:41:08.322989Z","iopub.status.idle":"2024-08-20T04:41:08.332985Z","shell.execute_reply.started":"2024-08-20T04:41:08.322967Z","shell.execute_reply":"2024-08-20T04:41:08.332192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"# from IPython.display import HTML\n# from base64 import b64encode\n\n# def play(filename):\n#     html = ''\n#     video = open(filename,'rb').read()\n#     src = 'data:video/mp4;base64,' + b64encode(video).decode()\n#     html += '<video width=1000 controls autoplay loop><source src=\"%s\" type=\"video/mp4\"></video>' % src \n#     return HTML(html)\n\n# play('/kaggle/working/exoplanet_signal.mp4')","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:08.626434Z","iopub.execute_input":"2024-08-20T04:41:08.627114Z","iopub.status.idle":"2024-08-20T04:41:08.63067Z","shell.execute_reply.started":"2024-08-20T04:41:08.62709Z","shell.execute_reply":"2024-08-20T04:41:08.629825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\ndef lecun_init(m):\n    \"\"\"\n    Apply LeCun initialization to the modules of the model.\n    \"\"\"\n    if isinstance(m, nn.Conv3d):\n        # LeCun initialization for Conv3d layers\n        nn.init.kaiming_uniform_(m.weight, a=5**0.5)\n        if m.bias is not None:\n            nn.init.constant_(m.bias, 0)\n    elif isinstance(m, nn.Linear):\n        # LeCun initialization for Linear layers\n        nn.init.kaiming_uniform_(m.weight, a=5**0.5)\n        if m.bias is not None:\n            nn.init.constant_(m.bias, 0)\n\nclass ResidualBlock3D(nn.Module):\n    def __init__(self, in_channels, out_channels, stride=1, drop_rate=0.1):\n        super(ResidualBlock3D, self).__init__()\n        self.conv1 = nn.Conv3d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1)\n        self.bn1 = nn.BatchNorm3d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n        self.dropout = nn.Dropout3d(p=drop_rate)  # Dropout after activation\n        self.conv2 = nn.Conv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)\n        self.bn2 = nn.BatchNorm3d(out_channels)\n\n        # Define the shortcut connection\n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=stride, padding=0),\n                nn.BatchNorm3d(out_channels)\n            )\n\n    def forward(self, x):\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.dropout(out)  # Apply dropout\n        out = self.bn2(self.conv2(out))\n        out += self.shortcut(x)\n        out = self.relu(out)\n        return out\n\nclass Residual3DCNN(nn.Module):\n    def __init__(self, drop_rate=0.1):\n        super(Residual3DCNN, self).__init__()\n\n        # Modify the number of input channels to match your data\n        self.layer1 = nn.Sequential(\n            ResidualBlock3D(1, 16, drop_rate=drop_rate),\n            ResidualBlock3D(16, 16, drop_rate=drop_rate)\n        )\n        self.layer2 = nn.Sequential(\n            ResidualBlock3D(16, 32, stride=2, drop_rate=drop_rate),\n            ResidualBlock3D(32, 32, drop_rate=drop_rate)\n        )\n        self.layer3 = nn.Sequential(\n            ResidualBlock3D(32, 64, stride=2, drop_rate=drop_rate),\n            ResidualBlock3D(64, 64, drop_rate=drop_rate)\n        )\n        \n        self.pool = nn.MaxPool3d(kernel_size=(2, 2, 2))\n        self.avgpool = nn.AdaptiveAvgPool3d((1, 1, 1))\n        self.dropout = nn.Dropout(p=drop_rate)  # Dropout before the final classification\n        self.fc1 = nn.Linear(64, 283)  # Adjust based on output size after flattening\n\n    def forward(self, x):\n        x = x.unsqueeze(1)\n        x = self.layer1(x)\n        x = self.pool(x)\n        x = self.layer2(x)\n        x = self.pool(x)\n        x = self.layer3(x)\n        x = self.pool(x)\n        \n        x = self.avgpool(x)\n        x = torch.flatten(x, 1)\n        x = self.dropout(x)  # Apply dropout\n        x = self.fc1(x)\n        \n        return x\n\nmodel = Residual3DCNN()\nmodel.apply(lecun_init)  # Apply LeCun initialization\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:09.165805Z","iopub.execute_input":"2024-08-20T04:41:09.166125Z","iopub.status.idle":"2024-08-20T04:41:09.372579Z","shell.execute_reply.started":"2024-08-20T04:41:09.166101Z","shell.execute_reply":"2024-08-20T04:41:09.371658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class Modified3DModel(nn.Module):\n#     def __init__(self, num_classes=283):\n#         super().__init__()\n\n#         # Load a pretrained model\n#         self.model = torch.hub.load('facebookresearch/pytorchvideo', 'x3d_l', pretrained=True)\n        \n#         # Modify the input channels of the first convolutional layer\n#         self.model.blocks[0].conv = nn.Conv3d(\n#             in_channels=1,  # Change to 1 channel\n#             out_channels=self.model.blocks[0].conv.conv_t.out_channels,\n#             kernel_size=self.model.blocks[0].conv.conv_t.kernel_size,\n#             stride=self.model.blocks[0].conv.conv_t.stride,\n#             padding=self.model.blocks[0].conv.conv_t.padding,\n#             bias=False,\n#         )\n        \n#         # Adapt the output layer to match the number of classes\n#         self.model.blocks[5].output_pool = nn.AdaptiveAvgPool3d(1)\n#         self.fc = nn.Linear(self.model.blocks[5].proj.out_features, num_classes)\n    \n#     def forward(self, x):\n#         x = self.model(x)\n#         x = self.fc(x)\n#         return x\n# num_classes = 283  # Example number of classes for your task\n# model = Modified3DModel(num_classes=num_classes)\n# model.to(device)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-08-20T04:41:09.374035Z","iopub.execute_input":"2024-08-20T04:41:09.374327Z","iopub.status.idle":"2024-08-20T04:41:09.379049Z","shell.execute_reply.started":"2024-08-20T04:41:09.374302Z","shell.execute_reply":"2024-08-20T04:41:09.378206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_indices = splits[fold][0]\nvalid_indices = splits[fold][1]\ntrain_ids = [all_planet_ids[i] for i in train_indices]\nvalid_ids = [all_planet_ids[i] for i in valid_indices]\ntrain_loader, valid_loader = create_dataloaders_for_fold(train_ids, valid_ids, instrument)\n# _ = next(iter(train_loader))\n# model(_['signal'].to(device)).shape","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:09.501425Z","iopub.execute_input":"2024-08-20T04:41:09.501714Z","iopub.status.idle":"2024-08-20T04:41:09.684182Z","shell.execute_reply.started":"2024-08-20T04:41:09.501691Z","shell.execute_reply":"2024-08-20T04:41:09.683249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.optim.lr_scheduler import OneCycleLR\n\nEPOCHS = config.EPOCHS\nBATCHES = len(train_loader)\nsteps = []\nlrs = []\noptim_lrs = []\nmodel = Residual3DCNN()\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)\nscheduler = OneCycleLR(\n    optimizer,\n    max_lr=1e-6,\n    epochs=config.EPOCHS,\n    steps_per_epoch=len(train_loader),\n    pct_start=0.1,\n    anneal_strategy=\"cos\",\n    final_div_factor=100,\n)\nfor epoch in range(EPOCHS):\n    for batch in range(BATCHES):\n        scheduler.step()\n        lrs.append(scheduler.get_last_lr()[0])\n        steps.append(epoch * BATCHES + batch)\n\nmax_lr = max(lrs)\nmin_lr = min(lrs)\nprint(f\"Maximum LR: {max_lr} | Minimum LR: {min_lr}\")\nplt.figure()\nplt.plot(steps, lrs, label='OneCycle')\nplt.ticklabel_format(axis='y', style='sci', scilimits=(0,0))\nplt.xlabel(\"Step\")\nplt.ylabel(\"Learning Rate\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:09.723893Z","iopub.execute_input":"2024-08-20T04:41:09.724201Z","iopub.status.idle":"2024-08-20T04:41:11.236431Z","shell.execute_reply.started":"2024-08-20T04:41:09.724176Z","shell.execute_reply":"2024-08-20T04:41:11.235585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn.functional as F\n\nclass R2Loss(nn.Module):\n    def __init__(self):\n        super(R2Loss, self).__init__()\n\n    def forward(self, y_pred, y_true):\n        # Calculate R-squared\n        ss_tot = torch.sum((y_true - torch.mean(y_true))**2)\n        ss_res = torch.sum((y_true - y_pred)**2)\n        r2 = 1 - (ss_res / ss_tot)\n        return -r2\n\n# Example usage with dummy data\ny_pred = F.softmax(torch.randn(2, 283, requires_grad=True),dim=-1)\ny_true = F.softmax(torch.randn(2, 283),dim=-1)\nloss = R2Loss()(y_pred, y_true)\nprint(f\"R² Loss: {loss.item()}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:11.238082Z","iopub.execute_input":"2024-08-20T04:41:11.238491Z","iopub.status.idle":"2024-08-20T04:41:11.273932Z","shell.execute_reply.started":"2024-08-20T04:41:11.238457Z","shell.execute_reply":"2024-08-20T04:41:11.273075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_epoch(train_loader, model, optimizer, epoch, scheduler, device):\n    \"\"\"One epoch training pass for Residual3DCNN.\"\"\"\n    model.train()\n    criterion = nn.MSELoss()\n    scaler = torch.cuda.amp.GradScaler(enabled=config.AMP)\n    losses = AverageMeter()\n    start = end = time.time()\n    global_step = 0\n    softmax = nn.Softmax(dim=-1)\n\n    # ========== ITERATE OVER TRAIN BATCHES ============\n    with tqdm(train_loader, unit=\"train_batch\", desc='Train') as tqdm_train_loader:\n        for step, batch in enumerate(tqdm_train_loader):\n            X = batch.pop(\"signal\").to(device)  # send inputs to `device`\n            y = batch.pop(\"label\").to(device)  # send labels to `device`\n            batch_size = y.size(0)\n            with torch.cuda.amp.autocast(enabled=config.AMP):\n                y_preds = model(X)\n                loss = criterion(y_preds.squeeze(), y.squeeze())\n            if config.GRADIENT_ACCUMULATION_STEPS > 1:\n                loss = loss / config.GRADIENT_ACCUMULATION_STEPS\n            losses.update(loss.item(), batch_size)\n            scaler.scale(loss).backward()\n            grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), config.MAX_GRAD_NORM)\n\n            if (step + 1) % config.GRADIENT_ACCUMULATION_STEPS == 0:\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n                global_step += 1\n                scheduler.step()\n            end = time.time()\n\n            # ========== LOG INFO ==========\n            #if step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                      'Elapsed {remain:s} '\n                      'Loss: {loss:.4f} '\n                      'Grad: {grad_norm:.4f}  '\n                      'LR: {lr:.8f}  '\n                      .format(epoch+1, step, len(train_loader),\n                              remain=timeSince(start, float(step+1)/len(train_loader)),\n                              loss=losses.avg*1e4,\n                              grad_norm=grad_norm,\n                              lr=scheduler.get_last_lr()[0]))\n\n    return losses.avg\ndef valid_epoch(valid_loader, model, device):\n    model.eval()\n    criterion = nn.MSELoss()  # Replace with your custom loss function\n    softmax = nn.Softmax(dim=-1)  # Adjust if needed\n    losses = AverageMeter()\n    prediction_dict = {}\n    preds = []\n    gts = []\n    start = end = time.time()\n    with tqdm(valid_loader, unit=\"valid_batch\", desc='Validation') as tqdm_valid_loader:\n        for step, batch in enumerate(tqdm_valid_loader):\n            X = batch.pop(\"signal\").to(device)\n            y = batch.pop(\"label\").to(device)\n            batch_size = y.size(0)\n            with torch.no_grad():\n                y_preds = model(X)\n                loss = criterion(y_preds.squeeze(), y.squeeze())\n            if config.GRADIENT_ACCUMULATION_STEPS > 1:\n                loss = loss / config.GRADIENT_ACCUMULATION_STEPS\n            losses.update(loss.item(), batch_size)\n            # Assuming `y_preds` are probabilities; adjust if they are raw logits or other formats\n            preds.append(y_preds.to('cpu').numpy())\n            gts.append(y.to('cpu').numpy())\n            end = time.time()\n\n            # ========== LOG INFO ==========\n            if step == (len(valid_loader)-1):\n                print('EVAL: [{0}/{1}] '\n                      'Elapsed {remain:s} '\n                      'Loss: {loss:.4f} '\n                      .format(step, len(valid_loader),\n                              remain=timeSince(start, float(step+1)/len(valid_loader)),\n                              loss=losses.avg*1e4))\n\n    prediction_dict[\"predictions\"] = np.concatenate(preds)\n    prediction_dict[\"targets\"] = np.concatenate(gts)\n    return losses.avg, prediction_dict\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:11.275351Z","iopub.execute_input":"2024-08-20T04:41:11.275707Z","iopub.status.idle":"2024-08-20T04:41:11.294045Z","shell.execute_reply.started":"2024-08-20T04:41:11.275677Z","shell.execute_reply":"2024-08-20T04:41:11.293244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_loop(fold):\n    LOGGER.info(f\"========== Fold: {fold} training ==========\")\n    train_indices = splits[fold][0]\n    valid_indices = splits[fold][1]\n    train_ids = [all_planet_ids[i] for i in train_indices]\n    valid_ids = [all_planet_ids[i] for i in valid_indices]\n    train_loader, valid_loader = create_dataloaders_for_fold(train_ids, valid_ids, instrument)\n\n    # ======== MODEL ==========\n    model = Residual3DCNN()\n    \n    model.to(device)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e6)\n    scheduler = OneCycleLR(\n        optimizer,\n        max_lr=1e-6,\n        epochs=config.EPOCHS,\n        steps_per_epoch=len(train_loader),\n        pct_start=0.1,\n        anneal_strategy=\"cos\",\n        final_div_factor=100,\n    )\n\n    # ======= LOSS ==========\n    criterion = criterion = R2Loss()\n\n    best_loss = np.inf\n    # ====== ITERATE EPOCHS ========\n    for epoch in range(config.EPOCHS):\n        start_time = time.time()\n\n        # ======= TRAIN ==========\n        avg_train_loss = train_epoch(train_loader, model, optimizer, epoch, scheduler, device)\n\n        # ======= EVALUATION ==========\n        avg_val_loss, prediction_dict = valid_epoch(valid_loader, model, device)\n        predictions = prediction_dict[\"predictions\"]\n        gts = prediction_dict[\"targets\"]\n\n        # ======= SCORING ==========\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_train_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n\n        if avg_val_loss < best_loss:\n            best_loss = avg_val_loss\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Loss: {best_loss:.4f} Model')\n            torch.save({'model': model.state_dict(),\n                        'predictions': predictions,\n                        'targets': gts\n                        },\n                         os.path.join(config.OUTPUT_DIR, f\"residual3dcnn_fold_{fold}_best.pth\")\n                         )\n\n    results = torch.load(os.path.join(config.OUTPUT_DIR, f\"residual3dcnn_fold_{fold}_best.pth\"),\n                             map_location=torch.device('cpu'))\n    valid_folds[target_preds] = results['predictions']\n    valid_folds[targets] = results['targets']\n\n    torch.cuda.empty_cache()\n    gc.collect()\n\n    return valid_folds\n# def inference_loop(df, fold):\n#     LOGGER.info(f\"========== Fold: {fold} inference ==========\")\n\n#     # ======== SPLIT ==========\n#     valid_folds = df[df['fold'] == fold].reset_index(drop=True)\n\n#     # ======== DATASETS ==========\n#     valid_dataset = CustomDataset(valid_folds, eeg_df, config, downsample=1, mode=\"test\")\n\n#     # ======== DATALOADERS ==========\n#     valid_loader = DataLoader(valid_dataset,\n#                               batch_size=config.BATCH_SIZE_VALID,\n#                               shuffle=False,\n#                               num_workers=config.NUM_WORKERS, pin_memory=True, drop_last=False)\n\n#     # ======== MODEL ==========\n#     model = Residual3DCNN(in_channels=1,  # Adjust as needed\n#                           num_classes=config.num_classes,\n#                           drop_rate=0.1)  # Add any other hyperparameters if necessary\n#     checkpoint = torch.load(os.path.join(config.OUTPUT_DIR, f\"residual3dcnn_fold_{fold}_best.pth\"), map_location=device)\n#     model.load_state_dict(checkpoint[\"model\"])\n#     model.to(device)\n\n#     # ======= LOSS ==========\n#     criterion = nn.KLDivLoss(reduction=\"batchmean\")\n\n#     avg_val_loss, prediction_dict = valid_epoch(valid_loader, model, device)\n#     predictions = prediction_dict[\"predictions\"]\n#     gts = prediction_dict[\"targets\"]\n\n#     # ======= SCORING ==========\n#     LOGGER.info(f'avg_val_loss: {avg_val_loss:.4f}')\n\n#     valid_folds[target_preds] = predictions\n#     valid_folds[targets] = gts\n\n#     torch.cuda.empty_cache()\n#     gc.collect()\n\n#     return valid_folds\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:11.295833Z","iopub.execute_input":"2024-08-20T04:41:11.296082Z","iopub.status.idle":"2024-08-20T04:41:11.309284Z","shell.execute_reply.started":"2024-08-20T04:41:11.296061Z","shell.execute_reply":"2024-08-20T04:41:11.308562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"oof_df = pd.DataFrame()\nfor fold in range(5):\n    if fold in [0]:\n        _oof_df = train_loop(fold)\n        oof_df = pd.concat([oof_df, _oof_df])\n        LOGGER.info(f\"========== Fold {fold} finished ==========\")\noof_df = oof_df.reset_index(drop=True)\noof_df.to_csv(os.path.join(config.OUTPUT_DIR, 'oof_df.csv'), index=False)","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:41:11.310296Z","iopub.execute_input":"2024-08-20T04:41:11.310571Z"},"trusted":true},"execution_count":null,"outputs":[]}]}