{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":71549,"databundleVersionId":8561470}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Ablation study layer 4 - Full setup","metadata":{}},{"cell_type":"code","source":"global LOCAL_AVAILABLE_SERIES\nLOCAL_AVAILABLE_SERIES = index_local_subset(SUBSET_ROOT)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T06:28:21.893513Z","iopub.execute_input":"2026-03-25T06:28:21.893859Z","iopub.status.idle":"2026-03-25T06:29:16.937493Z","shell.execute_reply.started":"2026-03-25T06:28:21.893822Z","shell.execute_reply":"2026-03-25T06:29:16.936842Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\ntorch.cuda.is_available()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T06:28:07.731923Z","iopub.execute_input":"2026-03-25T06:28:07.735414Z","iopub.status.idle":"2026-03-25T06:28:11.542679Z","shell.execute_reply.started":"2026-03-25T06:28:07.735371Z","shell.execute_reply":"2026-03-25T06:28:11.541848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nfrom torchvision.models.feature_extraction import create_feature_extractor\nimport sys\nimport random\nfrom sklearn.metrics import f1_score, classification_report\nimport cv2 # NEW: Import OpenCV for resizing\n\n# --- 1. Configuration Constants ---\nFEATURE_DIM = 768\nCLASSIFIER_INPUT_DIM = FEATURE_DIM * 2\nDISC_LEVEL_SLICE_COUNT = 7\nNUM_TARGETS = 25\nNUM_CLASSES = 3\nCOORD_PATH = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\"\nMAIN_LABEL_PATH = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\"\nSUBSET_ROOT = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\nCHECKPOINT_PATH = \"hybrid_model_checkpoint.pth\"\nTARGET_SLICE_SIZE = 256 # NEW: Standard resolution for resizing non-uniform slices\n\n# --- Training Loop Hyperparameters (Added for BTP planning) ---\nWORKING_BATCH_SIZE = 5   # Actual size of data processed per forward pass (Keep this low to avoid OOM)\nSIMULATED_BATCH_SIZE = 200 # Target batch size for stable training (2nd place equivalent)\nNUM_EPOCHS = 2\nLOG_INTERVAL = 1 # Report loss every 10 epochs\n\n# --- Optimization Parameters ---\nLEARNING_RATE = 1e-4\nWEIGHT_DECAY = 1e-5\nCLASS_WEIGHTS = torch.tensor([1.0, 2.0, 4.0], dtype=torch.float32) \n\n# --- NEW: Calculate Accumulation Steps ---\nACCUMULATION_STEPS = SIMULATED_BATCH_SIZE // WORKING_BATCH_SIZE\nif SIMULATED_BATCH_SIZE % WORKING_BATCH_SIZE != 0:\n    print(\"Warning: SIMULATED_BATCH_SIZE is not perfectly divisible by WORKING_BATCH_SIZE.\")\n    sys.exit(1) # Stop if not perfectly divisible for this prototype\n\n# --- Global Data Variables ---\ntry:\n    GLOBAL_COORDS_DF = pd.read_csv(COORD_PATH)\n    GLOBAL_MAIN_LABELS_DF = pd.read_csv(MAIN_LABEL_PATH)\nexcept FileNotFoundError:\n    print(f\"❌ CRITICAL ERROR: Cannot find required files: {COORD_PATH} or {MAIN_LABEL_PATH}\")\n    sys.exit(1)\n\nLOCAL_AVAILABLE_SERIES = [] \n\n# List of all 25 target columns (MUST match train.csv order)\nTARGET_COLUMNS = [\n    'spinal_canal_stenosis_l1_l2', 'left_neural_foraminal_narrowing_l1_l2', 'right_neural_foraminal_narrowing_l1_l2', 'left_subarticular_stenosis_l1_l2', 'right_subarticular_stenosis_l1_l2',\n    'spinal_canal_stenosis_l2_l3', 'left_neural_foraminal_narrowing_l2_l3', 'right_neural_foraminal_narrowing_l2_l3', 'left_subarticular_stenosis_l2_l3', 'right_subarticular_stenosis_l2_l3',\n    'spinal_canal_stenosis_l3_l4', 'left_neural_foraminal_narrowing_l3_l4', 'right_neural_foraminal_narrowing_l3_l4', 'left_subarticular_stenosis_l3_l4', 'right_subarticular_stenosis_l3_l4',\n    'spinal_canal_stenosis_l4_l5', 'left_neural_foraminal_narrowing_l4_l5', 'right_neural_foraminal_narrowing_l4_l5', 'left_subarticular_stenosis_l4_l5', 'right_subarticular_stenosis_l4_l5',\n    'spinal_canal_stenosis_l5_s1', 'left_neural_foraminal_narrowing_l5_s1', 'right_neural_foraminal_narrowing_l5_s1', 'left_subarticular_stenosis_l5_s1', 'right_subarticular_stenosis_l5_s1'\n]\nif len(TARGET_COLUMNS) != 25:\n    print(\"FATAL ERROR: Target column list size mismatch.\")\n    sys.exit(1)\n\n\n# --- Label Mapping ---\nLABEL_TO_INT = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T06:28:11.544076Z","iopub.execute_input":"2026-03-25T06:28:11.544489Z","iopub.status.idle":"2026-03-25T06:28:21.850068Z","shell.execute_reply.started":"2026-03-25T06:28:11.544463Z","shell.execute_reply":"2026-03-25T06:28:21.849452Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# --- 2. Model Architecture Components ---\n\nclass AttentionPooling(nn.Module):\n    def __init__(self, feature_dim=FEATURE_DIM):\n        super().__init__()\n        self.attention_weights = nn.Sequential(\n            nn.Linear(feature_dim, feature_dim // 4),\n            nn.ReLU(),\n            nn.Linear(feature_dim // 4, 1)\n        )\n    def forward(self, features):\n        scores = self.attention_weights(features) \n        weights = F.softmax(scores, dim=0)\n        F_aggregated = torch.sum(features * weights, dim=0)\n        return F_aggregated.unsqueeze(0), weights.squeeze(1)\n\nclass MultiHeadClassifier(nn.Module):\n    def __init__(self, input_dim=CLASSIFIER_INPUT_DIM):\n        super().__init__()\n        self.heads = nn.ModuleList([\n            nn.Linear(input_dim, NUM_CLASSES) for _ in range(NUM_TARGETS)\n        ])\n    def forward(self, x):\n        outputs = [head(x) for head in self.heads]\n        return outputs \n\n# --- 3. Data Utilities ---\n\ndef index_local_subset(root_dir):\n    \"\"\"Scans the local subset folder and creates a list of available (study, series) IDs.\"\"\"\n    available_series = []\n    for study_id_str in os.listdir(root_dir):\n        study_path = os.path.join(root_dir, study_id_str)\n        if os.path.isdir(study_path):\n            for series_id_str in os.listdir(study_path):\n                series_path = os.path.join(study_path, series_id_str)\n                if os.path.isdir(series_path):\n                    if any(f.endswith('.dcm') for f in os.listdir(series_path)):\n                        try:\n                            available_series.append((int(study_id_str), int(series_id_str)))\n                        except ValueError:\n                            continue\n                            \n    print(f\"✅ Indexed {len(available_series)} unique series locally available in '{root_dir}'.\")\n    return available_series\n\ndef load_and_preprocess_series(study_id, series_id):\n    \"\"\"Loads DICOMs, sorts, RESIZES to fixed shape, and normalizes.\"\"\"\n    series_path = os.path.join(SUBSET_ROOT, str(study_id), str(series_id))\n    \n    dicom_files = sorted([os.path.join(series_path, f) for f in os.listdir(series_path) if f.endswith('.dcm')])\n    if not dicom_files: return None, None\n\n    slices = []\n    for f in dicom_files:\n        try:\n            ds = pydicom.dcmread(f)\n            if hasattr(ds, 'InstanceNumber'): slices.append(ds)\n        except Exception: continue\n    slices.sort(key=lambda x: x.InstanceNumber)\n    \n    # NEW: Store processed slices after resizing\n    processed_images = []\n    for s in slices:\n        img_array = s.pixel_array.astype(np.float32)\n        \n        # Rescale the image array to the fixed target size (e.g., 256x256)\n        # This fixes the \"must have the same shape\" error\n        resized_img = cv2.resize(img_array, (TARGET_SLICE_SIZE, TARGET_SLICE_SIZE), \n                                 interpolation=cv2.INTER_LINEAR)\n        processed_images.append(resized_img)\n    \n    if not processed_images: return None, None\n    \n    image = np.stack(processed_images) # Now all arrays have the same shape\n    \n    # Normalized (0 to 1)\n    image = (image - image.min()) / (image.max() - image.min() + 1e-6)\n    \n    instance_map = {int(s.InstanceNumber): i for i, s in enumerate(slices)}\n    return torch.from_numpy(image).unsqueeze(1).float(), instance_map\n\ndef extract_coordinate_based_crop(full_series_tensor, instance_map, instance_number, x_coord, y_coord, crop_size=256):\n    \"\"\"Uses coordinate data to crop N=7 slices (Simulated YOLOX output).\"\"\"\n    S, C, H, W = full_series_tensor.shape\n    center_idx = instance_map.get(instance_number)\n    if center_idx is None: return None\n\n    # NOTE: Since the full_series_tensor is now TARGET_SLICE_SIZE x TARGET_SLICE_SIZE (256x256), \n    # the coordinate mapping is accurate for cropping.\n    \n    half_slices = DISC_LEVEL_SLICE_COUNT // 2\n    start_idx = max(0, center_idx - half_slices)\n    end_idx = min(S, center_idx + half_slices + 1)\n    \n    half_crop = crop_size // 2\n    x_min, x_max = max(0, int(x_coord) - half_crop), min(W, int(x_coord) + half_crop)\n    y_min, y_max = max(0, int(y_coord) - half_crop), min(H, int(y_coord) + half_crop)\n    \n    cropped_slices = full_series_tensor[start_idx:end_idx, :, y_min:y_max, x_min:x_max]\n    \n    pad_h = crop_size - cropped_slices.shape[2]\n    pad_w = crop_size - cropped_slices.shape[3]\n    if pad_h > 0 or pad_w > 0:\n        cropped_slices = F.pad(cropped_slices, (0, pad_w, 0, pad_h), 'constant', 0)\n        \n    return cropped_slices.repeat(1, 3, 1, 1)\n\ndef prepare_random_batch(batch_size=WORKING_BATCH_SIZE):\n    \"\"\"Selects a random batch of valid annotations and retrieves true labels.\"\"\"\n    \n    # 1. Get local annotations subset\n    available_series_df = pd.DataFrame(LOCAL_AVAILABLE_SERIES, columns=['study_id', 'series_id'])\n    local_annotations_df = pd.merge(GLOBAL_COORDS_DF, available_series_df, on=['study_id', 'series_id'])\n    \n    if local_annotations_df.empty:\n        raise Exception(\"No matching annotations found in local subset. Check file structure/names.\")\n        \n    n_sample = min(batch_size, len(local_annotations_df))\n    random_annotations = local_annotations_df.sample(n=n_sample, replace=False)\n\n    # 2. Extract true labels (25 elements per study)\n    annot_with_labels_df = pd.merge(\n        random_annotations, \n        GLOBAL_MAIN_LABELS_DF[['study_id'] + TARGET_COLUMNS], \n        on='study_id',\n        how='left' \n    ).dropna(subset=TARGET_COLUMNS) # Drop incomplete labels\n    \n    if annot_with_labels_df.empty:\n        raise Exception(\"Sampled annotations have incomplete labels and were dropped.\")\n\n    # 3. Convert string labels to integer tensor [Batch_Size, 25]\n    true_labels_int_list = []\n    \n    for _, row in annot_with_labels_df.iterrows():\n        string_labels = row[TARGET_COLUMNS].tolist()\n        try:\n            int_labels = [LABEL_TO_INT[label] for label in string_labels]\n            true_labels_int_list.append(int_labels)\n        except KeyError:\n            # print(f\"Warning: Skipping study {row['study_id']} due to unknown label value.\")\n            continue\n            \n    true_labels_tensor = torch.tensor(true_labels_int_list, dtype=torch.long)\n\n    # 4. Final Batch Data (Only include successfully labeled items)\n    batch_data = annot_with_labels_df.to_dict('records')\n    \n    if len(batch_data) != true_labels_tensor.shape[0]:\n        batch_data = batch_data[:true_labels_tensor.shape[0]] \n\n    return batch_data, true_labels_tensor\n\n# --- 4. Prediction and Loss Functions ---\n\ndef calculate_weighted_log_loss(raw_outputs, target_labels_batch):\n    \"\"\"Calculates the total Weighted Log Loss (Training Metric).\"\"\"\n    device = raw_outputs[0].device\n    loss_weights = CLASS_WEIGHTS.to(device)\n    loss_fn = nn.CrossEntropyLoss(weight=loss_weights)\n    \n    total_loss = 0\n    for target_idx in range(NUM_TARGETS):\n        logits = raw_outputs[target_idx]\n        targets = target_labels_batch[:, target_idx]\n        total_loss += loss_fn(logits, targets)\n        \n    return total_loss\n\ndef calculate_accuracy(predicted_classes, true_classes):\n    \"\"\"Calculates overall categorical accuracy.\"\"\"\n    correct_predictions = (predicted_classes == true_classes).sum()\n    total_predictions = len(true_classes)\n    return correct_predictions.item() / total_predictions\n\ndef get_final_probabilities(raw_outputs):\n    \"\"\"Converts raw logit outputs into the final 25*3 probability vector (Inference).\"\"\"\n    all_probabilities = [F.softmax(logits, dim=1) for logits in raw_outputs]\n    final_predictions = torch.cat(all_probabilities, dim=1)\n    return final_predictions\n\ndef save_model_checkpoint(cnn_model, attention_pooler, classifier_heads, optimizer, path):\n    \"\"\"Saves the state dictionaries of all trainable components.\"\"\"\n    state = {\n        'cnn_state_dict': cnn_model.state_dict(),\n        'attn_pooler_state_dict': attention_pooler.state_dict(),\n        'classifier_heads_state_dict': classifier_heads.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'feature_dim': FEATURE_DIM,\n        'targets': NUM_TARGETS\n    }\n    torch.save(state, path)\n    print(f\"\\n✅ Model checkpoint saved successfully to: {path}\")\n\n# --- 5. Main Execution Function ---\nimport gc\ndef run_training_evaluation_pass(batch_size=WORKING_BATCH_SIZE):\n    gc.collect()\n    torch.cuda.empty_cache()\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n\n    global LOCAL_AVAILABLE_SERIES\n    LOCAL_AVAILABLE_SERIES = index_local_subset(SUBSET_ROOT)\n\n    if not LOCAL_AVAILABLE_SERIES:\n        print(\"❌ CRITICAL ERROR: No local DICOM series found. Cannot run training.\")\n        return\n\n    # 1. Initialize Models\n    cnn_model = timm.create_model('convnext_tiny', pretrained=True).to(device)\n    cnn_feature_extractor = create_feature_extractor(cnn_model, return_nodes={'head.flatten': 'features'}).to(device)\n    attention_pooler = AttentionPooling(FEATURE_DIM).to(device)\n    classifier_heads = MultiHeadClassifier().to(device)\n    \n    # Setup Optimizer\n    all_trainable_params = list(cnn_feature_extractor.parameters()) + \\\n                           list(attention_pooler.parameters()) + \\\n                           list(classifier_heads.parameters())\n    optimizer = torch.optim.AdamW(all_trainable_params, lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)\n\n    # Global tracking for average loss across epochs\n    epoch_losses = []\n    \n    # --- START OF EPOCH LOOP (Simulated Full Training) ---\n    print(f\"\\n--- Starting Simulated Training for {NUM_EPOCHS} Epochs for {SIMULATED_BATCH_SIZE} series---\")\n    \n    for epoch in range(NUM_EPOCHS):\n        \n        # Set models to training mode\n        cnn_feature_extractor.train()\n        attention_pooler.train()\n        classifier_heads.train()\n        \n        # Reset total loss for this accumulation cycle\n        accumulated_loss = 0\n        successful_samples_in_cycle = 0\n        \n        # --- ACCUMULATION LOOP: SIMULATES PROCESSING A LARGE BATCH ---\n        for accumulation_step in range(ACCUMULATION_STEPS):\n            \n            try:\n                batch_items, true_labels_tensor = prepare_random_batch(WORKING_BATCH_SIZE)\n                if not batch_items: continue\n            except Exception as e:\n                # print(f\"❌ Batch preparation failed in epoch {epoch}, step {accumulation_step}: {e}\")\n                continue\n\n            # Containers for the working batch data\n            fused_features_list = []\n            \n            # 2. Forward Pass: Feature Extraction & Fusion\n            for i, item in enumerate(batch_items):\n                full_series_tensor, instance_map = load_and_preprocess_series(item['study_id'], item['series_id'])\n                if full_series_tensor is None: continue \n\n                input_slices = extract_coordinate_based_crop(\n                    full_series_tensor, instance_map, item['instance_number'], item['x'], item['y']\n                ).to(device)\n                if input_slices is None: continue\n\n                # CNN Feature Extraction, ViT/Attention Aggregation, and Fusion\n                slice_features = [cnn_feature_extractor(slice_img.unsqueeze(0))['features'].squeeze()\n                                  for slice_img in input_slices]\n                F_sequence = torch.stack(slice_features)\n                F_aggregated, _ = attention_pooler(F_sequence)\n                F_central = F_sequence[DISC_LEVEL_SLICE_COUNT // 2].unsqueeze(0)\n                F_fused = torch.cat([F_aggregated, F_central], dim=1)\n                \n                fused_features_list.append(F_fused)\n\n            if not fused_features_list: continue \n\n            F_batch = torch.cat(fused_features_list, dim=0) \n            batch_actual_size = F_batch.shape[0]\n\n            # Trim true labels to match the actual batch size\n            true_labels_tensor_trimmed = true_labels_tensor[:batch_actual_size].to(device)\n            \n            # A. Forward Pass (Classification)\n            raw_outputs = classifier_heads(F_batch.to(device))\n            \n            # B. Calculate Loss (Scaled by accumulation steps)\n            total_loss = calculate_weighted_log_loss(raw_outputs, true_labels_tensor_trimmed)\n            \n            # C. BACKWARD PASS (accumulate gradients)\n            # We normalize the loss here to prevent gradients from exploding\n            total_loss = total_loss / ACCUMULATION_STEPS\n            total_loss.backward()\n            \n            accumulated_loss += total_loss.item() * ACCUMULATION_STEPS # Store unscaled loss for logging\n            successful_samples_in_cycle += batch_actual_size\n\n        # --- END ACCUMULATION LOOP ---\n        \n        if successful_samples_in_cycle > 0:\n            # D. OPTIMIZER STEP (Update weights once per simulated batch)\n            optimizer.step()\n            \n            # E. Zero the gradients for the next cycle\n            optimizer.zero_grad() \n            \n            # Store average loss for reporting\n            epoch_losses.append(accumulated_loss / successful_samples_in_cycle) # Store average loss per sample\n\n            # --- Report Loss Interval (Every 10 Epochs) ---\n            if (epoch + 1) % LOG_INTERVAL == 0:\n                avg_loss_log = sum(epoch_losses[-LOG_INTERVAL:]) / LOG_INTERVAL\n                print(f\"Epoch {epoch + 1:02d}/{NUM_EPOCHS} | Accumulated Loss): {avg_loss_log:.4f}\")\n        else:\n            # Re-zero gradients if cycle failed completely\n            optimizer.zero_grad() \n            print(f\"Epoch {epoch + 1:02d}/{NUM_EPOCHS} | Warning: No samples processed in this epoch.\")\n            \n    # --- END OF EPOCH LOOP ---\n    \n    # --- 4. Final Evaluation and Reporting (After ALL Epochs) ---\n    \n    # Set models to evaluation mode for final metrics\n    cnn_feature_extractor.eval()\n    attention_pooler.eval()\n    classifier_heads.eval()\n    \n    # Run ONE final random batch for robust metric reporting\n    try:\n        final_batch_items, true_labels_tensor_final = prepare_random_batch(SIMULATED_BATCH_SIZE)\n        if not final_batch_items: raise Exception(\"Final batch preparation failed.\")\n    except Exception as e:\n        print(f\"\\n❌ FINAL EVALUATION FAILED: Could not prepare validation batch. Error: {e}\")\n        return\n\n    # Process the final validation batch\n    # NOTE: Since feature extraction is outside the main loop, we run it again for the final batch\n    final_fused_features = []\n    with torch.no_grad():\n        for item in final_batch_items:\n            full_series_tensor, instance_map = load_and_preprocess_series(item['study_id'], item['series_id'])\n            if full_series_tensor is None: continue \n\n            input_slices = extract_coordinate_based_crop(\n                full_series_tensor, instance_map, item['instance_number'], item['x'], item['y']\n            ).to(device)\n            if input_slices is None: continue\n\n            # CNN Feature Extraction, ViT/Attention Aggregation, and Fusion\n            slice_features = [cnn_feature_extractor(slice_img.unsqueeze(0))['features'].squeeze()\n                              for slice_img in input_slices]\n            F_sequence = torch.stack(slice_features)\n            F_aggregated, _ = attention_pooler(F_sequence)\n            F_central = F_sequence[DISC_LEVEL_SLICE_COUNT // 2].unsqueeze(0)\n            F_fused = torch.cat([F_aggregated, F_central], dim=1)\n            \n            final_fused_features.append(F_fused)\n\n    F_batch_final = torch.cat(final_fused_features, dim=0)\n\n    raw_outputs_final = classifier_heads(F_batch_final)\n    \n    # Save the updated model state\n    # save_model_checkpoint(cnn_feature_extractor, attention_pooler, classifier_heads, optimizer, CHECKPOINT_PATH)\n    \n    # Get final hard predictions and reshape true labels\n    final_probs = get_final_probabilities(raw_outputs_final)\n    predicted_probs_reshaped = final_probs.view(-1, NUM_CLASSES)\n    predicted_classes = torch.argmax(predicted_probs_reshaped, dim=1).cpu().numpy()\n    \n    # Trim true labels to match final batch size (if any were skipped during processing)\n    true_labels_tensor_final = true_labels_tensor_final[:F_batch_final.shape[0]].to(device)\n    true_classes = true_labels_tensor_final.view(-1).cpu().numpy()\n    \n    # Calculate Metrics\n    f1 = f1_score(true_classes, predicted_classes, average='weighted', zero_division=0)\n    accuracy = calculate_accuracy(predicted_classes, true_classes)\n    \n    # Get target names for the report\n    target_names_map = {0: 'Mild', 1: 'Moderate', 2: 'Severe'}\n    target_names = [target_names_map[i] for i in range(NUM_CLASSES)]\n\n    # --- Report Results ---\n    print(\"\\n\" + \"=\"*75)\n    print(\"FINAL EVALUATION REPORT (After Simulated Training)\")\n    \n    print(\"\\n--- 1. TRAINING SUMMARY ---\")\n    print(f\"Total Epochs Simulated: {NUM_EPOCHS}\")\n    if epoch_losses:\n        print(f\"Final Average Training Loss: {epoch_losses[-1]:.4f}\")\n    print(f\"✅ Model weights saved to {CHECKPOINT_PATH}.\")\n    \n    print(\"\\n--- 2. PERFORMANCE METRICS (Final Validation Batch) ---\")\n    print(f\"Overall Categorical Accuracy: {accuracy:.4f}\")\n    print(f\"Weighted F1-Score: {f1:.4f}\")\n    \n    print(\"\\n--- CLASSIFICATION REPORT (Final Validation Batch) ---\")\n    print(classification_report(true_classes, predicted_classes, target_names=target_names, zero_division=0))\n    print(\"=\"*75)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T06:28:21.851469Z","iopub.execute_input":"2026-03-25T06:28:21.851812Z","iopub.status.idle":"2026-03-25T06:28:21.892610Z","shell.execute_reply.started":"2026-03-25T06:28:21.851761Z","shell.execute_reply":"2026-03-25T06:28:21.891880Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"WORKING_BATCH_SIZE = 10   # Actual size of data processed per forward pass\nSIMULATED_BATCH_SIZE = 1600 # Target batch size for stable training\nNUM_EPOCHS = 10\n\nimport time\nstart=time.time()\nrun_training_evaluation_pass()\nend = time.time()\nelapsed_time_seconds = end-start\nprint(f\"\\nTotal Elapsed Time: {elapsed_time_seconds:.2f} seconds\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"WORKING_BATCH_SIZE = 10   # Actual size of data processed per forward pass\nSIMULATED_BATCH_SIZE = 6294 # Target batch size for stable training\nNUM_EPOCHS = 45\n\nimport time\nstart=time.time()\nrun_training_evaluation_pass()\nend = time.time()\nelapsed_time_seconds = end-start\nprint(f\"\\nTotal Elapsed Time: {elapsed_time_seconds:.2f} seconds\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T06:29:47.269581Z","iopub.execute_input":"2026-03-25T06:29:47.269926Z","iopub.status.idle":"2026-03-25T12:41:17.384380Z","shell.execute_reply.started":"2026-03-25T06:29:47.269898Z","shell.execute_reply":"2026-03-25T12:41:17.383344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix\nimport numpy as np\n\ndef plot_final_confusion_matrix(y_true, y_pred):\n    labels = ['Mild', 'Moderate', 'Severe']\n    cm = confusion_matrix(y_true, y_pred)\n    \n    # Normalize by row (percentage of true labels)\n    cm_perc = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]\n    \n    plt.figure(figsize=(8, 6))\n    sns.heatmap(cm_perc, annot=True, fmt='.2%', cmap='Blues', \n                xticklabels=labels, yticklabels=labels)\n    plt.title('Normalized Confusion Matrix: Hybrid Model')\n    plt.ylabel('Actual Condition')\n    plt.xlabel('Predicted Condition')\n    plt.show()\n\n# You can call this using your validation results\nplot_final_confusion_matrix(true_classes, predicted_classes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T13:01:38.337872Z","iopub.execute_input":"2026-03-25T13:01:38.338638Z","iopub.status.idle":"2026-03-25T13:01:38.620841Z","shell.execute_reply.started":"2026-03-25T13:01:38.338609Z","shell.execute_reply":"2026-03-25T13:01:38.619845Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}