{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 2,
   "id": "5abbe242-38c9-444a-8a5e-403ed4db7581",
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "C:\\Users\\hm\\anaconda3\\envs\\image\\Lib\\site-packages\\requests\\__init__.py:86: RequestsDependencyWarning: Unable to find acceptable character detection dependency (chardet or charset_normalizer).\n",
      "  warnings.warn(\n",
      "C:\\Users\\hm\\anaconda3\\envs\\image\\Lib\\site-packages\\tqdm\\auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html\n",
      "  from .autonotebook import tqdm as notebook_tqdm\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "All libraries imported successfully!\n"
     ]
    }
   ],
   "source": [
    "# =====================\n",
    "# 1. IMPORTS AND SETUP\n",
    "# =====================\n",
    "\n",
    "import pandas as pd\n",
    "import numpy as np\n",
    "import torch\n",
    "import torch.nn as nn\n",
    "import torch.nn.functional as F\n",
    "from torch.utils.data import Dataset, DataLoader\n",
    "from sklearn.preprocessing import StandardScaler, LabelEncoder\n",
    "from sklearn.model_selection import train_test_split\n",
    "from sklearn.metrics import accuracy_score, roc_auc_score, f1_score, classification_report, confusion_matrix\n",
    "from sklearn.utils.class_weight import compute_class_weight\n",
    "from torch.optim.lr_scheduler import ReduceLROnPlateau\n",
    "from transformers import AutoTokenizer, AutoModel\n",
    "import re\n",
    "import warnings\n",
    "import joblib\n",
    "\n",
    "warnings.filterwarnings('ignore')\n",
    "\n",
    "print(\"All libraries imported successfully!\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "id": "93b6b406-f7bd-4722-8909-9ca85ad2b798",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Dataset loaded successfully!\n",
      "Dataset shape: (792, 18)\n",
      "Columns: ['label', 'HOSPITAL_EXPIRE_FLAG', 'ETHNICITY', 'AGE', 'GENDER', 'Heart_Rate', 'Systolic_BP', 'Diastolic_BP', 'Respiratory_Rate', 'SpO2', 'Temperature', 'Sodium', 'Potassium', 'Creatinine', 'BUN', 'Hemoglobin', 'WBC', 'text']\n"
     ]
    }
   ],
   "source": [
    "# =====================\n",
    "# 2. DATA LOADING\n",
    "# =====================\n",
    "\n",
    "# Load your dataset\n",
    "df = pd.read_csv('ICUDataset.csv')  # Change to your actual file path\n",
    "\n",
    "print(\"Dataset loaded successfully!\")\n",
    "print(f\"Dataset shape: {df.shape}\")\n",
    "print(f\"Columns: {df.columns.tolist()}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "id": "b76d1281-9111-444b-8786-e4ae7d1e944d",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\n",
      "=== TARGET COLUMN PROCESSING ===\n",
      "Original target distribution:\n",
      "label\n",
      "ELECTIVE     429\n",
      "EMERGENCY    363\n",
      "Name: count, dtype: int64\n",
      "Processed target distribution:\n",
      "label\n",
      "0    429\n",
      "1    363\n",
      "Name: count, dtype: int64\n"
     ]
    }
   ],
   "source": [
    "# =====================\n",
    "# 3. TARGET COLUMN PROCESSING\n",
    "# =====================\n",
    "\n",
    "# Your target column is 'label'\n",
    "target_column = 'label'\n",
    "\n",
    "if target_column not in df.columns:\n",
    "    raise ValueError(f\"Target column '{target_column}' not found!\")\n",
    "\n",
    "print(f\"\\n=== TARGET COLUMN PROCESSING ===\")\n",
    "print(f\"Original target distribution:\\n{df[target_column].value_counts()}\")\n",
    "\n",
    "# Process target column\n",
    "def process_target(df, target_col):\n",
    "    original_dtype = df[target_col].dtype\n",
    "    \n",
    "    if df[target_col].dtype == 'object':\n",
    "        label_mapping = {'EMERGENCY': 1, 'ELECTIVE': 0}\n",
    "        df[target_col] = df[target_col].map(label_mapping)\n",
    "        \n",
    "        if df[target_col].isna().any():\n",
    "            df[target_col] = pd.to_numeric(df[target_col], errors='coerce')\n",
    "    \n",
    "    # Convert to integer\n",
    "    df[target_col] = df[target_col].astype(float).fillna(0).astype(int)\n",
    "    \n",
    "    # Drop rows with invalid targets if any\n",
    "    if df[target_col].isna().any():\n",
    "        df = df.dropna(subset=[target_col])\n",
    "    \n",
    "    return df\n",
    "\n",
    "df = process_target(df, target_column)\n",
    "labels = df[target_column].values\n",
    "\n",
    "print(f\"Processed target distribution:\\n{df[target_column].value_counts()}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 8,
   "id": "cad7b0b1-8458-46a6-87fa-f951777a3b7b",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\n",
      "=== STRUCTURED DATA PROCESSING ===\n",
      "Available structured columns: ['HOSPITAL_EXPIRE_FLAG', 'ETHNICITY', 'AGE', 'GENDER', 'Heart_Rate', 'Systolic_BP', 'Diastolic_BP', 'Respiratory_Rate', 'SpO2', 'Temperature', 'Sodium', 'Potassium', 'Creatinine', 'BUN', 'Hemoglobin', 'WBC']\n",
      "Categorical: ['HOSPITAL_EXPIRE_FLAG', 'ETHNICITY', 'GENDER']\n",
      "Numerical: ['AGE', 'Heart_Rate', 'Systolic_BP', 'Diastolic_BP', 'Respiratory_Rate', 'SpO2', 'Temperature', 'Sodium', 'Potassium', 'Creatinine', 'BUN', 'Hemoglobin', 'WBC']\n",
      "Structured features shape: (792, 16)\n"
     ]
    }
   ],
   "source": [
    "# =====================\n",
    "# 4. STRUCTURED DATA PROCESSING\n",
    "# =====================\n",
    "\n",
    "print(\"\\n=== STRUCTURED DATA PROCESSING ===\")\n",
    "\n",
    "# Define feature columns\n",
    "structured_columns = [\n",
    "    'HOSPITAL_EXPIRE_FLAG', 'ETHNICITY', 'AGE', 'GENDER', \n",
    "    'Heart_Rate', 'Systolic_BP', 'Diastolic_BP', 'Respiratory_Rate', \n",
    "    'SpO2', 'Temperature', 'Sodium', 'Potassium', 'Creatinine', \n",
    "    'BUN', 'Hemoglobin', 'WBC'\n",
    "]\n",
    "\n",
    "# Get only available columns\n",
    "available_columns = [col for col in structured_columns if col in df.columns]\n",
    "print(f\"Available structured columns: {available_columns}\")\n",
    "\n",
    "# Separate categorical and numerical\n",
    "categorical_columns = [col for col in ['HOSPITAL_EXPIRE_FLAG', 'ETHNICITY', 'GENDER'] if col in available_columns]\n",
    "numerical_columns = [col for col in available_columns if col not in categorical_columns]\n",
    "\n",
    "print(f\"Categorical: {categorical_columns}\")\n",
    "print(f\"Numerical: {numerical_columns}\")\n",
    "\n",
    "# Handle missing values\n",
    "for col in available_columns:\n",
    "    if col in categorical_columns:\n",
    "        df[col].fillna(df[col].mode()[0] if not df[col].mode().empty else 'Unknown', inplace=True)\n",
    "    else:\n",
    "        df[col].fillna(df[col].median(), inplace=True)\n",
    "\n",
    "# Encode categorical variables\n",
    "label_encoders = {}\n",
    "for col in categorical_columns:\n",
    "    le = LabelEncoder()\n",
    "    df[col] = le.fit_transform(df[col].astype(str))\n",
    "    label_encoders[col] = le\n",
    "\n",
    "# Normalize numerical features\n",
    "if numerical_columns:\n",
    "    scaler = StandardScaler()\n",
    "    df[numerical_columns] = scaler.fit_transform(df[numerical_columns])\n",
    "\n",
    "# Create structured features\n",
    "structured_features = df[available_columns].values.astype(np.float32)\n",
    "print(f\"Structured features shape: {structured_features.shape}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 10,
   "id": "0031f8cd-0a8f-4044-a0ec-fe2d02e11030",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\n",
      "=== TEXT DATA PROCESSING ===\n",
      "Generating text embeddings...\n",
      "Processed 16/792 samples\n",
      "Processed 176/792 samples\n",
      "Processed 336/792 samples\n",
      "Processed 496/792 samples\n",
      "Processed 656/792 samples\n",
      "Text embeddings shape: (792, 768)\n"
     ]
    }
   ],
   "source": [
    "\n",
    "# =====================\n",
    "# 5. TEXT DATA PROCESSING\n",
    "# =====================\n",
    "\n",
    "print(\"\\n=== TEXT DATA PROCESSING ===\")\n",
    "\n",
    "# Text cleaning function\n",
    "def clean_clinical_text(text):\n",
    "    if pd.isna(text) or text == \"\":\n",
    "        return \"no clinical notes available\"\n",
    "    \n",
    "    text = str(text).lower().strip()\n",
    "    \n",
    "    # Basic cleaning\n",
    "    text = re.sub(r'[^\\w\\s.,;:!?()%+-/]', ' ', text)\n",
    "    text = re.sub(r'\\s+', ' ', text)\n",
    "    \n",
    "    # Keep meaningful words\n",
    "    words = text.split()\n",
    "    medical_abbr = {'bp', 'hr', 'rr', 'spo2', 'temp', 'wbc', 'bun', 'na', 'k', 'hgb', 'hct'}\n",
    "    words = [word for word in words if len(word) > 2 or word in medical_abbr]\n",
    "    \n",
    "    cleaned_text = ' '.join(words)\n",
    "    return cleaned_text if cleaned_text else \"clinical assessment documented\"\n",
    "\n",
    "# Clean text\n",
    "df['text_clean'] = df['text'].apply(clean_clinical_text)\n",
    "\n",
    "\n",
    "# Text embedding class\n",
    "class ClinicalTextEmbedder:\n",
    "    def __init__(self, model_name=\"emilyalsentzer/Bio_ClinicalBERT\"):\n",
    "        self.tokenizer = AutoTokenizer.from_pretrained(model_name)\n",
    "        self.model = AutoModel.from_pretrained(model_name)\n",
    "        self.model.eval()\n",
    "    \n",
    "    def embed_texts(self, texts, max_length=256):\n",
    "        embeddings = []\n",
    "        for i, text in enumerate(texts):\n",
    "            try:\n",
    "                inputs = self.tokenizer(\n",
    "                    text, \n",
    "                    max_length=max_length, \n",
    "                    padding='max_length', \n",
    "                    truncation=True, \n",
    "                    return_tensors='pt'\n",
    "                )\n",
    "                with torch.no_grad():\n",
    "                    outputs = self.model(**inputs)\n",
    "                    embedding = outputs.last_hidden_state[:, 0, :].squeeze()\n",
    "                embeddings.append(embedding.numpy())\n",
    "            except Exception as e:\n",
    "                print(f\"Error with text {i}: {e}\")\n",
    "                embeddings.append(np.zeros(768))  # Default ClinicalBERT size\n",
    "        return np.array(embeddings)\n",
    "\n",
    "# Generate text embeddings\n",
    "print(\"Generating text embeddings...\")\n",
    "embedder = ClinicalTextEmbedder()\n",
    "\n",
    "# Process in batches to avoid memory issues\n",
    "batch_size = 16\n",
    "text_embeddings_list = []\n",
    "\n",
    "for i in range(0, len(df), batch_size):\n",
    "    batch_texts = df['text_clean'].iloc[i:i+batch_size].tolist()\n",
    "    batch_embeddings = embedder.embed_texts(batch_texts)\n",
    "    text_embeddings_list.append(batch_embeddings)\n",
    "    \n",
    "    if (i // batch_size) % 10 == 0:\n",
    "        print(f\"Processed {min(i + batch_size, len(df))}/{len(df)} samples\")\n",
    "\n",
    "text_embeddings = np.vstack(text_embeddings_list)\n",
    "print(f\"Text embeddings shape: {text_embeddings.shape}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 11,
   "id": "bf28fc43-886c-4d09-b251-7e6ecbb1222c",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\n",
      "=== CREATING DATASET ===\n",
      "Training samples: 633\n",
      "Testing samples: 159\n"
     ]
    }
   ],
   "source": [
    "\n",
    "\n",
    "\n",
    "# =====================\n",
    "# 6. CREATE DATASET AND DATALOADERS\n",
    "# =====================\n",
    "\n",
    "print(\"\\n=== CREATING DATASET ===\")\n",
    "\n",
    "class MultimodalICUDataset(Dataset):\n",
    "    def __init__(self, structured_features, text_embeddings, labels):\n",
    "        self.structured_features = torch.FloatTensor(structured_features)\n",
    "        self.text_embeddings = torch.FloatTensor(text_embeddings)\n",
    "        self.labels = torch.LongTensor(labels)\n",
    "    \n",
    "    def __len__(self):\n",
    "        return len(self.labels)\n",
    "    \n",
    "    def __getitem__(self, idx):\n",
    "        return {\n",
    "            'structured': self.structured_features[idx],\n",
    "            'text': self.text_embeddings[idx],\n",
    "            'label': self.labels[idx]\n",
    "        }\n",
    "\n",
    "# Create full dataset\n",
    "full_dataset = MultimodalICUDataset(structured_features, text_embeddings, labels)\n",
    "\n",
    "# Split into train/test\n",
    "train_idx, test_idx = train_test_split(\n",
    "    range(len(full_dataset)), \n",
    "    test_size=0.2, \n",
    "    stratify=labels,\n",
    "    random_state=42\n",
    ")\n",
    "\n",
    "train_dataset = torch.utils.data.Subset(full_dataset, train_idx)\n",
    "test_dataset = torch.utils.data.Subset(full_dataset, test_idx)\n",
    "\n",
    "# Create data loaders\n",
    "train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)\n",
    "test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False)\n",
    "\n",
    "print(f\"Training samples: {len(train_dataset)}\")\n",
    "print(f\"Testing samples: {len(test_dataset)}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 16,
   "id": "4e9a0854-ce31-4203-83a3-fcb6ae6db2e9",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\n",
      "=== BUILDING ENHANCED MODEL ===\n",
      "Model created with 310,210 parameters\n",
      "\n",
      "Testing model with sample batch...\n",
      "Sample output shape: torch.Size([32, 2])\n",
      "✓ Model architecture is correct!\n"
     ]
    }
   ],
   "source": [
    "\n",
    "# =====================\n",
    "# 7. FIXED ENHANCED MODEL ARCHITECTURE\n",
    "# =====================\n",
    "\n",
    "print(\"\\n=== BUILDING ENHANCED MODEL ===\")\n",
    "\n",
    "class EnhancedMultimodalModel(nn.Module):\n",
    "    def __init__(self, structured_dim, text_embedding_dim, hidden_dim=256, num_classes=2, dropout_rate=0.5):\n",
    "        super(EnhancedMultimodalModel, self).__init__()\n",
    "        \n",
    "        # Structured data pathway\n",
    "        self.structured_net = nn.Sequential(\n",
    "            nn.Linear(structured_dim, hidden_dim),\n",
    "            nn.BatchNorm1d(hidden_dim),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(dropout_rate),\n",
    "            nn.Linear(hidden_dim, hidden_dim // 2),\n",
    "            nn.BatchNorm1d(hidden_dim // 2),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(dropout_rate)\n",
    "        )\n",
    "        \n",
    "        # Text data pathway\n",
    "        self.text_net = nn.Sequential(\n",
    "            nn.Linear(text_embedding_dim, hidden_dim),\n",
    "            nn.BatchNorm1d(hidden_dim),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(dropout_rate),\n",
    "            nn.Linear(hidden_dim, hidden_dim // 2),\n",
    "            nn.BatchNorm1d(hidden_dim // 2),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(dropout_rate)\n",
    "        )\n",
    "        \n",
    "        # Feature fusion - FIXED: Now concatenates instead of adding\n",
    "        self.fusion = nn.Sequential(\n",
    "            nn.Linear(hidden_dim, hidden_dim // 2),  # hidden_dim from structured + text pathways\n",
    "            nn.BatchNorm1d(hidden_dim // 2),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(dropout_rate)\n",
    "        )\n",
    "        \n",
    "        # Final classifier\n",
    "        self.classifier = nn.Sequential(\n",
    "            nn.Linear(hidden_dim // 2, hidden_dim // 4),\n",
    "            nn.BatchNorm1d(hidden_dim // 4),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(dropout_rate // 2),\n",
    "            nn.Linear(hidden_dim // 4, num_classes)\n",
    "        )\n",
    "        \n",
    "    def forward(self, structured, text):\n",
    "        # Process both modalities\n",
    "        structured_out = self.structured_net(structured)  # shape: [batch_size, hidden_dim//2]\n",
    "        text_out = self.text_net(text)  # shape: [batch_size, hidden_dim//2]\n",
    "        \n",
    "        # FIXED: Concatenate features instead of adding\n",
    "        combined = torch.cat([structured_out, text_out], dim=1)  # shape: [batch_size, hidden_dim]\n",
    "        \n",
    "        # Fusion layer\n",
    "        fused = self.fusion(combined)  # shape: [batch_size, hidden_dim//2]\n",
    "        \n",
    "        # Final classification\n",
    "        output = self.classifier(fused)\n",
    "        return output\n",
    "\n",
    "# Initialize model\n",
    "model = EnhancedMultimodalModel(\n",
    "    structured_dim=structured_features.shape[1],\n",
    "    text_embedding_dim=text_embeddings.shape[1],\n",
    "    hidden_dim=256,\n",
    "    dropout_rate=0.5\n",
    ")\n",
    "\n",
    "print(f\"Model created with {sum(p.numel() for p in model.parameters()):,} parameters\")\n",
    "\n",
    "# Test the model with a sample batch to ensure it works\n",
    "print(\"\\nTesting model with sample batch...\")\n",
    "sample_batch = next(iter(train_loader))\n",
    "sample_structured = sample_batch['structured']\n",
    "sample_text = sample_batch['text']\n",
    "try:\n",
    "    sample_output = model(sample_structured, sample_text)\n",
    "    print(f\"Sample output shape: {sample_output.shape}\")\n",
    "    print(\"✓ Model architecture is correct!\")\n",
    "except Exception as e:\n",
    "    print(f\"✗ Model error: {e}\")\n",
    "    # Try alternative simpler model\n",
    "    print(\"Trying alternative model...\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 18,
   "id": "4a0eab93-ba7a-4841-a6b6-2d7345dbe8f7",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\n",
      "=== TRAINING SETUP ===\n",
      "Class weights: tensor([0.9231, 1.0909])\n"
     ]
    }
   ],
   "source": [
    "\n",
    "# =====================\n",
    "# 8. ADVANCED TRAINING SETUP\n",
    "# =====================\n",
    "\n",
    "print(\"\\n=== TRAINING SETUP ===\")\n",
    "\n",
    "# Handle class imbalance\n",
    "class_weights = compute_class_weight(\n",
    "    'balanced',\n",
    "    classes=np.unique(labels),\n",
    "    y=labels\n",
    ")\n",
    "class_weights = torch.FloatTensor(class_weights)\n",
    "print(f\"Class weights: {class_weights}\")\n",
    "\n",
    "# Training components\n",
    "criterion = nn.CrossEntropyLoss(weight=class_weights)\n",
    "optimizer = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.01)\n",
    "scheduler = ReduceLROnPlateau(optimizer, mode='max', patience=5, factor=0.5)\n",
    "\n",
    "# Early stopping\n",
    "class EarlyStopping:\n",
    "    def __init__(self, patience=10, verbose=False, delta=0):\n",
    "        self.patience = patience\n",
    "        self.verbose = verbose\n",
    "        self.counter = 0\n",
    "        self.best_score = None\n",
    "        self.early_stop = False\n",
    "        self.delta = delta\n",
    "\n",
    "    def __call__(self, val_auc, model):\n",
    "        score = val_auc\n",
    "        if self.best_score is None:\n",
    "            self.best_score = score\n",
    "        elif score < self.best_score + self.delta:\n",
    "            self.counter += 1\n",
    "            if self.verbose:\n",
    "                print(f'EarlyStopping counter: {self.counter}/{self.patience}')\n",
    "            if self.counter >= self.patience:\n",
    "                self.early_stop = True\n",
    "        else:\n",
    "            self.best_score = score\n",
    "            self.counter = 0\n",
    "\n",
    "early_stopping = EarlyStopping(patience=10, verbose=True)\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 20,
   "id": "95324d8f-0e00-4d9e-9370-2baf713e9ce0",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\n",
      "=== STARTING TRAINING ===\n",
      "Current learning rate: 0.001000\n",
      "Epoch   0: Loss: 0.6812, Acc: 0.7484, AUC: 0.8372, F1: 0.7561\n",
      "New best model saved with AUC: 0.8372\n",
      "New best model saved with AUC: 0.8692\n",
      "New best model saved with AUC: 0.8793\n",
      "New best model saved with AUC: 0.8961\n",
      "New best model saved with AUC: 0.8965\n",
      "New best model saved with AUC: 0.9124\n",
      "EarlyStopping counter: 1/10\n",
      "EarlyStopping counter: 2/10\n",
      "EarlyStopping counter: 3/10\n",
      "EarlyStopping counter: 4/10\n",
      "Current learning rate: 0.001000\n",
      "EarlyStopping counter: 5/10\n",
      "Epoch  10: Loss: 0.1755, Acc: 0.7987, AUC: 0.8942, F1: 0.8025\n",
      "EarlyStopping counter: 6/10\n",
      "EarlyStopping counter: 7/10\n",
      "EarlyStopping counter: 8/10\n",
      "EarlyStopping counter: 9/10\n",
      "EarlyStopping counter: 10/10\n",
      "Early stopping at epoch 15\n",
      "\n",
      "*** BEST AUC: 0.9124 ***\n"
     ]
    }
   ],
   "source": [
    "# =====================\n",
    "# 9. TRAINING LOOP\n",
    "# =====================\n",
    "\n",
    "print(\"\\n=== STARTING TRAINING ===\")\n",
    "\n",
    "def train_model(model, train_loader, test_loader, epochs=100):\n",
    "    train_losses = []\n",
    "    val_metrics = []\n",
    "    best_auc = 0\n",
    "    \n",
    "    for epoch in range(epochs):\n",
    "        # Training\n",
    "        model.train()\n",
    "        epoch_loss = 0\n",
    "        \n",
    "        for batch in train_loader:\n",
    "            structured = batch['structured']\n",
    "            text = batch['text']\n",
    "            batch_labels = batch['label']\n",
    "            \n",
    "            optimizer.zero_grad()\n",
    "            outputs = model(structured, text)\n",
    "            loss = criterion(outputs, batch_labels)\n",
    "            loss.backward()\n",
    "            optimizer.step()\n",
    "            \n",
    "            epoch_loss += loss.item()\n",
    "        \n",
    "        # Validation\n",
    "        model.eval()\n",
    "        all_preds = []\n",
    "        all_labels = []\n",
    "        all_probs = []\n",
    "        \n",
    "        with torch.no_grad():\n",
    "            for batch in test_loader:\n",
    "                structured = batch['structured']\n",
    "                text = batch['text']\n",
    "                batch_labels = batch['label']\n",
    "                \n",
    "                outputs = model(structured, text)\n",
    "                probs = F.softmax(outputs, dim=1)\n",
    "                preds = torch.argmax(outputs, dim=1)\n",
    "                \n",
    "                all_preds.extend(preds.cpu().numpy())\n",
    "                all_labels.extend(batch_labels.cpu().numpy())\n",
    "                all_probs.extend(probs[:, 1].cpu().numpy())\n",
    "        \n",
    "        # Calculate metrics\n",
    "        accuracy = accuracy_score(all_labels, all_preds)\n",
    "        auc = roc_auc_score(all_labels, all_probs)\n",
    "        f1 = f1_score(all_labels, all_preds)\n",
    "        \n",
    "        train_losses.append(epoch_loss / len(train_loader))\n",
    "        val_metrics.append({'accuracy': accuracy, 'auc': auc, 'f1': f1})\n",
    "        \n",
    "        # Update scheduler\n",
    "        scheduler.step(auc)\n",
    "        \n",
    "        # Print learning rate every 10 epochs\n",
    "        if epoch % 10 == 0:\n",
    "            current_lr = optimizer.param_groups[0]['lr']\n",
    "            print(f\"Current learning rate: {current_lr:.6f}\")\n",
    "        \n",
    "        # Early stopping check\n",
    "        early_stopping(auc, model)\n",
    "        \n",
    "        # Print progress\n",
    "        if epoch % 10 == 0 or epoch == epochs - 1:\n",
    "            print(f'Epoch {epoch:3d}: Loss: {epoch_loss/len(train_loader):.4f}, '\n",
    "                  f'Acc: {accuracy:.4f}, AUC: {auc:.4f}, F1: {f1:.4f}')\n",
    "        \n",
    "        # Save best model\n",
    "        if auc > best_auc:\n",
    "            best_auc = auc\n",
    "            torch.save(model.state_dict(), 'best_icu_model.pth')\n",
    "            print(f\"New best model saved with AUC: {auc:.4f}\")\n",
    "        \n",
    "        # Check early stopping\n",
    "        if early_stopping.early_stop:\n",
    "            print(f\"Early stopping at epoch {epoch}\")\n",
    "            break\n",
    "    \n",
    "    print(f\"\\n*** BEST AUC: {best_auc:.4f} ***\")\n",
    "    return train_losses, val_metrics\n",
    "\n",
    "# Start training\n",
    "train_losses, val_metrics = train_model(model, train_loader, test_loader, epochs=100)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 22,
   "id": "80f9cab8-ac04-4597-b286-20ebf4881519",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\n",
      "=== FINAL EVALUATION ===\n",
      "\n",
      "*** FINAL RESULTS ***\n",
      "Accuracy: 0.8302\n",
      "AUC: 0.9124\n",
      "F1-Score: 0.8138\n",
      "\n",
      "Classification Report:\n",
      "              precision    recall  f1-score   support\n",
      "\n",
      "           0       0.84      0.85      0.84        86\n",
      "           1       0.82      0.81      0.81        73\n",
      "\n",
      "    accuracy                           0.83       159\n",
      "   macro avg       0.83      0.83      0.83       159\n",
      "weighted avg       0.83      0.83      0.83       159\n",
      "\n",
      "\n",
      "Confusion Matrix:\n",
      "[[73 13]\n",
      " [14 59]]\n"
     ]
    }
   ],
   "source": [
    "\n",
    "# =====================\n",
    "# 10. EVALUATION\n",
    "# =====================\n",
    "\n",
    "print(\"\\n=== FINAL EVALUATION ===\")\n",
    "\n",
    "# Load best model\n",
    "model.load_state_dict(torch.load('best_icu_model.pth'))\n",
    "model.eval()\n",
    "\n",
    "# Final predictions\n",
    "all_preds = []\n",
    "all_labels = []\n",
    "all_probs = []\n",
    "\n",
    "with torch.no_grad():\n",
    "    for batch in test_loader:\n",
    "        structured = batch['structured']\n",
    "        text = batch['text']\n",
    "        batch_labels = batch['label']\n",
    "        \n",
    "        outputs = model(structured, text)\n",
    "        probs = F.softmax(outputs, dim=1)\n",
    "        preds = torch.argmax(outputs, dim=1)\n",
    "        \n",
    "        all_preds.extend(preds.cpu().numpy())\n",
    "        all_labels.extend(batch_labels.cpu().numpy())\n",
    "        all_probs.extend(probs[:, 1].cpu().numpy())\n",
    "\n",
    "# Final metrics\n",
    "final_accuracy = accuracy_score(all_labels, all_preds)\n",
    "final_auc = roc_auc_score(all_labels, all_probs)\n",
    "final_f1 = f1_score(all_labels, all_preds)\n",
    "\n",
    "print(\"\\n*** FINAL RESULTS ***\")\n",
    "print(f\"Accuracy: {final_accuracy:.4f}\")\n",
    "print(f\"AUC: {final_auc:.4f}\")\n",
    "print(f\"F1-Score: {final_f1:.4f}\")\n",
    "\n",
    "print(\"\\nClassification Report:\")\n",
    "print(classification_report(all_labels, all_preds))\n",
    "\n",
    "print(\"\\nConfusion Matrix:\")\n",
    "print(confusion_matrix(all_labels, all_preds))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fc35a01a-493f-4e4f-a6de-1a74c10a185a",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": 24,
   "id": "e58216eb-2312-40fc-83cb-46f43b1e15e0",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\n",
      "=== SAVING PIPELINE ===\n",
      "Pipeline saved successfully!\n",
      "Files saved: final_icu_model.pth, scaler.pkl, label_encoders.pkl\n",
      "\n",
      "=== COMPLETE! ===\n"
     ]
    }
   ],
   "source": [
    "\n",
    "# =====================\n",
    "# 11. SAVE COMPLETE PIPELINE\n",
    "# =====================\n",
    "\n",
    "print(\"\\n=== SAVING PIPELINE ===\")\n",
    "\n",
    "# Save model\n",
    "torch.save(model.state_dict(), 'final_icu_model.pth')\n",
    "\n",
    "# Save preprocessing objects\n",
    "joblib.dump(scaler, 'scaler.pkl')\n",
    "joblib.dump(label_encoders, 'label_encoders.pkl')\n",
    "\n",
    "print(\"Pipeline saved successfully!\")\n",
    "print(\"Files saved: final_icu_model.pth, scaler.pkl, label_encoders.pkl\")\n",
    "\n",
    "print(\"\\n=== COMPLETE! ===\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ea71873f-2eab-4d32-8e5a-3853645fa9a2",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a62b12ca-9d37-4f28-b306-6940dbda7b03",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "785de550-8cfe-40d1-a961-b5539ca763c5",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fc9d2600-dd54-4dd3-ada4-873cefb97543",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python [conda env:image]",
   "language": "python",
   "name": "conda-env-image-py"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.12.11"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
