{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":23249,"databundleVersionId":2399555,"sourceType":"competition"}],"dockerImageVersionId":31240,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Environment, Configuration, and Model Setup","metadata":{}},{"cell_type":"code","source":"# G2Net Gravitational Wave Detection\n# Final Project Presentation\n# Architecture: EfficientNet-B2 + CQT (Bandpass Filtered)\n\nimport os\nimport random\nimport time\nimport datetime\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score, confusion_matrix, roc_curve\nfrom scipy import signal\nfrom tqdm.notebook import tqdm\n\n# Suppress warnings for clean output\nwarnings.simplefilter(\"ignore\")\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\n\n# Install necessary libraries silently\nprint(\"[INFO] Installing libraries...\")\nos.system('pip install -q nnAudio timm > /dev/null 2>&1')\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.cuda.amp as amp\nimport timm\nfrom nnAudio.Spectrogram import CQT1992v2\n\n# -------------------------------------------------------------------------\n# CONFIGURATION\n# -------------------------------------------------------------------------\nclass CFG:\n    seed = 42\n    model_name = 'tf_efficientnet_b2_ns'\n    img_size = 256\n    batch_size = 64\n    epochs = 8 \n    subset_size = 60000\n    lr = 1e-3\n    weight_decay = 1e-4\n    n_fold = 5\n    # CQT Parameters\n    sr = 2048\n    fmin = 20\n    fmax = 1024\n    hop_length = 32\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    num_workers = 2\n\n# -------------------------------------------------------------------------\n# UTILS\n# -------------------------------------------------------------------------\ndef set_seed(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\ndef get_file_path(image_id):\n    return \"../input/g2net-gravitational-wave-detection/train/{}/{}/{}/{}.npy\".format(\n        image_id[0], image_id[1], image_id[2], image_id\n    )\n\ndef format_time(elapsed):\n    return str(datetime.timedelta(seconds=int(round((elapsed)))))\n\n# -------------------------------------------------------------------------\n# PREPROCESSING\n# -------------------------------------------------------------------------\ndef apply_bandpass(x, lf=20, hf=500, order=4, sr=2048):\n    \"\"\"Whitening filter to remove low-frequency noise floor.\"\"\"\n    sos = signal.butter(order, [lf, hf], btype=\"bandpass\", output=\"sos\", fs=sr)\n    normalization = np.sqrt(1/2048)\n    return signal.sosfiltfilt(sos, x) * normalization\n\n# -------------------------------------------------------------------------\n# DATASET\n# -------------------------------------------------------------------------\nclass G2NetDataset(Dataset):\n    def __init__(self, df):\n        self.df = df\n        self.file_names = df['id'].values\n        self.labels = df['target'].values\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_path = get_file_path(self.file_names[idx])\n        waves = np.load(file_path).astype(np.float32)\n        \n        # Bandpass Filter\n        for i in range(3):\n            waves[i] = apply_bandpass(waves[i])\n            \n        # Normalization\n        waves = waves / np.max(np.abs(waves), axis=1, keepdims=True)\n        \n        return torch.tensor(waves, dtype=torch.float32), torch.tensor(self.labels[idx], dtype=torch.float32)\n\n# -------------------------------------------------------------------------\n# MODEL\n# -------------------------------------------------------------------------\nclass G2NetModel(nn.Module):\n    def __init__(self, cfg, pretrained=True):\n        super(G2NetModel, self).__init__()\n        self.cqt = CQT1992v2(\n            sr=cfg.sr, fmin=cfg.fmin, fmax=cfg.fmax,\n            hop_length=cfg.hop_length,\n            output_format=\"Magnitude\", verbose=False\n        )\n        self.backbone = timm.create_model(\n            cfg.model_name, pretrained=pretrained, in_chans=3,\n            num_classes=1, drop_rate=0.3, drop_path_rate=0.2\n        )\n\n    def forward(self, x):\n        bs, ch, time_dim = x.shape\n        x = x.view(bs * ch, time_dim)\n        x = self.cqt(x)\n        x = torch.log1p(x)\n        x = x.view(bs, ch, x.size(1), x.size(2))\n        \n        # Standardize\n        mean = x.mean(dim=(2, 3), keepdim=True)\n        std = x.std(dim=(2, 3), keepdim=True)\n        x = (x - mean) / (std + 1e-7)\n        \n        if x.shape[2] != CFG.img_size:\n            x = torch.nn.functional.interpolate(x, size=(CFG.img_size, CFG.img_size), \n                                              mode='bilinear', align_corners=False)\n        return self.backbone(x)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-23T01:19:59.544048Z","iopub.execute_input":"2025-12-23T01:19:59.544339Z","iopub.status.idle":"2025-12-23T01:21:29.342472Z","shell.execute_reply.started":"2025-12-23T01:19:59.544309Z","shell.execute_reply":"2025-12-23T01:21:29.341875Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loop","metadata":{}},{"cell_type":"code","source":"# -------------------------------------------------------------------------\n# TRAINING LOOPS\n# -------------------------------------------------------------------------\ndef train_epoch(train_loader, model, criterion, optimizer, scaler, device):\n    model.train()\n    running_loss = 0\n    preds_list = []\n    labels_list = []\n    \n    pbar = tqdm(train_loader, desc=\"Training\", leave=False, bar_format='{l_bar}{bar:10}{r_bar}')\n    \n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        \n        with amp.autocast():\n            outputs = model(images).squeeze(1)\n            loss = criterion(outputs, labels)\n            \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        running_loss += loss.item()\n        preds_list.append(torch.sigmoid(outputs).detach().cpu().numpy())\n        labels_list.append(labels.detach().cpu().numpy())\n        \n        pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n        \n    all_preds = np.concatenate(preds_list)\n    all_labels = np.concatenate(labels_list)\n    \n    return running_loss/len(train_loader), roc_auc_score(all_labels, all_preds)\n\ndef validate_epoch(valid_loader, model, criterion, device):\n    model.eval()\n    running_loss = 0\n    preds_list = []\n    labels_list = []\n    \n    with torch.no_grad():\n        for images, labels in tqdm(valid_loader, desc=\"Validation\", leave=False):\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images).squeeze(1)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            preds_list.append(torch.sigmoid(outputs).cpu().numpy())\n            labels_list.append(labels.cpu().numpy())\n            \n    all_preds = np.concatenate(preds_list)\n    all_labels = np.concatenate(labels_list)\n    \n    return running_loss/len(valid_loader), roc_auc_score(all_labels, all_preds), all_preds, all_labels\n\n# -------------------------------------------------------------------------\n# MAIN EXECUTION\n# -------------------------------------------------------------------------\nif __name__ == '__main__':\n    set_seed(CFG.seed)\n    print(f\"[INFO] V6 FINAL: EfficientNet-B2 | Epochs: {CFG.epochs} | Data: {CFG.subset_size}\")\n    \n    df = pd.read_csv(\"../input/g2net-gravitational-wave-detection/training_labels.csv\")\n    df_subset = df.sample(n=CFG.subset_size, random_state=CFG.seed).reset_index(drop=True)\n    \n    skf = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\n    \n    oof_df = df_subset.copy()\n    oof_df['pred_b2'] = 0.0\n    \n    # List to store results for all folds (to be used in the next cell)\n    history_data = [] \n\n    for fold, (train_idx, val_idx) in enumerate(skf.split(df_subset, df_subset['target'])):\n        print(f\"\\n=== FOLD {fold+1}/{CFG.n_fold} ===\")\n        \n        train_ds = G2NetDataset(df_subset.iloc[train_idx].reset_index(drop=True))\n        valid_ds = G2NetDataset(df_subset.iloc[val_idx].reset_index(drop=True))\n        \n        train_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True, \n                                num_workers=CFG.num_workers, pin_memory=True)\n        valid_loader = DataLoader(valid_ds, batch_size=CFG.batch_size, shuffle=False, \n                                num_workers=CFG.num_workers, pin_memory=True)\n        \n        model = G2NetModel(CFG).to(CFG.device)\n        optimizer = optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n        scaler = amp.GradScaler()\n        scheduler = optim.lr_scheduler.OneCycleLR(optimizer, max_lr=CFG.lr, \n                                                steps_per_epoch=len(train_loader), epochs=CFG.epochs)\n        \n        # Local dictionary for this fold\n        fold_history = {'t_loss': [], 'v_loss': [], 't_auc': [], 'v_auc': []}\n        \n        best_auc = 0\n        best_preds = None\n        \n        for epoch in range(CFG.epochs):\n            t_loss, t_auc = train_epoch(train_loader, model, nn.BCEWithLogitsLoss(), \n                                      optimizer, scaler, CFG.device)\n            v_loss, v_auc, v_preds, _ = validate_epoch(valid_loader, model, \n                                                     nn.BCEWithLogitsLoss(), CFG.device)\n            \n            # Record metrics\n            fold_history['t_loss'].append(t_loss)\n            fold_history['v_loss'].append(v_loss)\n            fold_history['t_auc'].append(t_auc)\n            fold_history['v_auc'].append(v_auc)\n            \n            # Print exact log format\n            print(f\"Ep {epoch+1}/{CFG.epochs} | Loss: {t_loss:.4f}/{v_loss:.4f} | AUC: {t_auc:.4f}/{v_auc:.4f}\")\n            \n            if v_auc > best_auc:\n                best_auc = v_auc\n                best_preds = v_preds\n            \n            scheduler.step()\n            \n        # Store predictions\n        oof_df.loc[val_idx, 'pred_b2'] = best_preds\n        \n        # Store history for the Final Visualization Cell\n        history_data.append(fold_history)\n\n    print(f\"\\n=== FINAL RESULTS ===\")\n    overall_auc = roc_auc_score(df_subset['target'], oof_df['pred_b2'])\n    print(f\"Global AUC (EfficientNet-B2): {overall_auc:.5f}\")\n    \n    # Save OOF for analysis\n    oof_df.to_csv('oof_predictions_b2.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T01:22:09.654562Z","iopub.execute_input":"2025-12-23T01:22:09.655200Z","iopub.status.idle":"2025-12-23T01:28:09.137525Z","shell.execute_reply.started":"2025-12-23T01:22:09.655171Z","shell.execute_reply":"2025-12-23T01:28:09.136454Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Final Project Visualization","metadata":{}},{"cell_type":"code","source":"# -------------------------------------------------------------------------\n# FINAL VISUALIZATION & REPORTING \n# -------------------------------------------------------------------------\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import roc_curve, roc_auc_score, confusion_matrix\n\n# Configuration for plots\nsns.set_style(\"white\")\nplt.rcParams['axes.spines.right'] = False\nplt.rcParams['axes.spines.top'] = False\nplt.rcParams['axes.grid'] = False\n\n# -------------------------------------------------------------------------\n# 1. RESULTS PER FOLD (GRID LAYOUT) \n# -------------------------------------------------------------------------\ndef plot_grid_results(histories):\n    n_folds = len(histories)\n    epochs = range(1, len(histories[0]['t_loss']) + 1)\n    colors = sns.color_palette(\"husl\", n_folds)\n    \n    fig, axes = plt.subplots(2, n_folds, figsize=(20, 8), sharex=True)\n    fig.text(0.125, 1.02, 'Results per fold', fontsize=28, fontweight='bold', ha='left')\n    \n    for i in range(n_folds):\n        h = histories[i]\n        c = colors[i]\n        \n        # ROW 1: LOSS\n        ax_l = axes[0, i]\n        ax_l.plot(epochs, h['t_loss'], label='Train', color=c, marker='o', lw=2)\n        ax_l.plot(epochs, h['v_loss'], label='Valid', color='gray', marker='x', ls='--', lw=1.5)\n        ax_l.set_title(f'Loss: Fold {i+1}', fontsize=14, fontweight='bold')\n        if i == 0: ax_l.set_ylabel('BCE Loss', fontsize=12)\n        ax_l.legend(loc='upper right')\n        \n        # ROW 2: AUC\n        ax_a = axes[1, i]\n        ax_a.plot(epochs, h['t_auc'], label='Train AUC', color=c, marker='o', lw=2)\n        ax_a.plot(epochs, h['v_auc'], label='Valid AUC', color='dimgray', marker='s', lw=2)\n        if i == 0: ax_a.set_ylabel('AUC Score', fontsize=12)\n        ax_a.set_xlabel('Epoch', fontsize=12)\n        ax_a.legend(loc='lower right')\n        \n    plt.tight_layout()\n    plt.show()\n\n# -------------------------------------------------------------------------\n# 2. HD SPECTROGRAMS \n# -------------------------------------------------------------------------\ndef plot_candidates_fixed_layout_inferno(dataset, device):\n    # Specific Candidates\n    target_ids = ['9c13c328bf', '3329ef4849']\n    print(f\"\\n[INFO] Generating Plot: Logica 'Hard Threshold' + Layout 'Clean'...\")\n    \n    # 1. Find indices\n    found_indices = []\n    for target in target_ids:\n        for i in range(len(dataset)):\n            if dataset.file_names[i] == target:\n                found_indices.append(i)\n                break\n                \n    if len(found_indices) < 2:\n        print(\"! Error: Targets not found in current subset.\")\n        return\n\n    # 2. Configure CQT for Visualization\n    # fmin=20, fmax=500, bins=12 for \"blocky\" definition\n    cqt_layer = CQT1992v2(sr=2048, fmin=20, fmax=500, hop_length=32, \n                          bins_per_octave=12, output_format=\"Magnitude\", \n                          verbose=False).to(device)\n\n    # 3. Plot Setup\n    fig, axes = plt.subplots(1, 2, figsize=(20, 7))\n    plt.subplots_adjust(right=0.9, top=0.85)\n\n    for i, idx in enumerate(found_indices):\n        wave, label = dataset[idx]\n        file_id = dataset.file_names[idx]\n        wave_tensor = wave.unsqueeze(0).to(device) # Shape: [1, 3, 4096]\n\n        with torch.no_grad():\n            # Flatten channels to process CQT\n            spec = cqt_layer(wave_tensor.view(1*3, -1))\n            spec_np = spec.cpu().numpy()\n            \n            # Combine H1 + L1 detectors\n            img = spec_np[0] + spec_np[1]\n\n            # Processing\n            img = np.log1p(img)\n            \n            # Median Subtraction\n            median_per_row = np.median(img, axis=1, keepdims=True)\n            img = img - median_per_row\n            \n            # Normalization 0-1\n            img = (img - np.min(img)) / (np.max(img) - np.min(img) + 1e-7)\n            \n            # HARD THRESHOLD 0.60\n            img[img < 0.60] = 0\n\n            # Plotting\n            ax = axes[i]\n            im = ax.imshow(img, aspect='auto', origin='lower', cmap='inferno', interpolation='bicubic')\n            \n            # Titles & Axes\n            ax.set_title(f\"ID: {file_id}\\nSource: LIGO Hanford + Livingston (Combined)\", \n                         fontsize=16, fontweight='bold', pad=10)\n            ax.set_ylabel(\"Frequency (Low to High)\", fontsize=12)\n            ax.set_xlabel(\"Time steps\", fontsize=12)\n            ax.grid(False)\n            ax.set_xticks([])\n            ax.set_yticks([])\n\n    # Colorbar\n    cbar_ax = fig.add_axes([0.92, 0.15, 0.02, 0.7])\n    cbar = fig.colorbar(im, cax=cbar_ax)\n    cbar.set_label('Normalized Amplitude (Threshold > 0.60)', fontsize=12)\n\n    plt.suptitle(\"Gravitational Wave Signal Reconstruction (Deep Learning)\", fontsize=22, y=1.05)\n    plt.show()\n\n# -------------------------------------------------------------------------\n# 3. GLOBAL ROC CURVE (SHADED) \n# -------------------------------------------------------------------------\ndef plot_final_roc_shaded(df):\n    y_true = df['target']\n    y_pred = df['pred_b2']\n    \n    fpr, tpr, _ = roc_curve(y_true, y_pred)\n    score = roc_auc_score(y_true, y_pred)\n    \n    plt.figure(figsize=(9, 7))\n    plt.plot(fpr, tpr, color='#8B0000', lw=3, label=f'ROC Curve (AUC = {score:.4f})')\n    plt.fill_between(fpr, tpr, color='#8B0000', alpha=0.1, label='AUC Area')\n    plt.plot([0, 1], [0, 1], color='navy', lw=1, linestyle='--')\n    plt.xlim([-0.01, 1.0])\n    plt.ylim([0.0, 1.05])\n    plt.xlabel('False Positive Rate', fontsize=14)\n    plt.ylabel('True Positive Rate', fontsize=14)\n    plt.title('Receiver Operating Characteristic (Global)', fontsize=20, fontweight='bold', pad=20)\n    plt.legend(loc=\"lower right\", fontsize=14)\n    sns.despine()\n    plt.show()\n\n# -------------------------------------------------------------------------\n# 4. GLOBAL CONFUSION MATRIX (FIXED LAYOUT)\n# -------------------------------------------------------------------------\ndef plot_global_confusion_matrix_fixed(df):\n    y_true = df['target']\n    y_pred_prob = df['pred_b2']\n    \n    # Threshold 0.5\n    y_pred_binary = (y_pred_prob > 0.5).astype(int)\n    cm = confusion_matrix(y_true, y_pred_binary)\n    \n    tn, fp, fn, tp = cm.ravel()\n    accuracy = (tp + tn) / len(y_true)\n    sensitivity = tp / (tp + fn)\n    specificity = tn / (tn + fp)\n    \n    group_names = ['True Neg', 'False Pos', 'False Neg', 'True Pos']\n    group_counts = [\"{0:0.0f}\".format(value) for value in cm.flatten()]\n    group_percentages = [\"{0:.2%}\".format(value) for value in cm.flatten()/np.sum(cm)]\n    \n    labels = [f\"{v1}\\n{v2}\\n{v3}\" for v1, v2, v3 in zip(group_names, group_counts, group_percentages)]\n    labels = np.asarray(labels).reshape(2,2)\n    \n    fig, ax = plt.subplots(figsize=(8, 7.5))\n    sns.set_style(\"white\")\n    sns.heatmap(cm, annot=labels, fmt='', cmap='Blues', cbar=False, \n                annot_kws={\"size\": 14, \"weight\": \"bold\"}, linewidths=2, linecolor='white')\n    \n    plt.title('Global Confusion Matrix (Threshold 0.5)', fontsize=18, fontweight='bold', pad=20)\n    plt.xlabel('Predicted Label', fontsize=14)\n    plt.ylabel('True Label', fontsize=14)\n    \n    # Bottom spacing\n    plt.subplots_adjust(bottom=0.25)\n    \n    stats_text = (f\"Accuracy: {accuracy:.4f}\\n\"\n                  f\"Sensitivity (Recall): {sensitivity:.4f}\\n\"\n                  f\"Specificity: {specificity:.4f}\")\n    \n    plt.figtext(0.5, 0.08, stats_text, ha=\"center\", fontsize=14, \n                bbox={\"facecolor\":\"orange\", \"alpha\":0.1, \"pad\":10, \"edgecolor\":\"orange\"})\n    plt.show()\n\n# -------------------------------------------------------------------------\n# EXECUTE VISUALIZATIONS\n# -------------------------------------------------------------------------\nif 'history_data' in globals() and 'valid_ds' in globals() and 'oof_df' in globals():\n    print(\"Generating Grid Results per Fold...\")\n    plot_grid_results(history_data)\n    \n    plot_candidates_fixed_layout_inferno(valid_ds, CFG.device)\n    \n    print(\"\\nGenerating ROC Curve...\")\n    plot_final_roc_shaded(oof_df)\n    \n    print(\"\\nGenerating Global Confusion Matrix...\")\n    plot_global_confusion_matrix_fixed(oof_df)\nelse:\n    print(\"! Variables not found. Please run the training cell first.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}