OSCR

Predicting ICU in-hospital mortality from text-encoded structured EHR data using adaptive transformer layer fusion.

Code ↔ Paper

13 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 13 matches
  1. [1] § STAR★Methods › Method details › The ALF module operates in several stages › Layer weight computation ↔ train.py, lines 115–182 · score 0.71 · global context vectors, LayerNorm, Context Fusion, concatenated, fused, weight
  2. [2] § STAR★Methods › Method details › The ALF module operates in several stages › Layer weight computation ↔ train.py, lines 185–251 · score 0.70 · layer attention mechanism, Attention scores, hidden states, attention mask, stacked, dimensions
  3. [3] § STAR★Methods › Method details › Attentional classifier head (ACH) module ↔ train.py, lines 305–339 · score 0.67 · expecting sequence, Attentional Classifier Head, fused embedding, unsqueezed, module
  4. [4] § STAR★Methods › Quantification and statistical analysis › Performance metrics ↔ inference.py, lines 712–771 · score 0.62 · confidence intervals, F1 score, Precision, Recall, bootstrap, metric
  5. [5] § STAR★Methods › Method details › The ALF module operates in several stages › Layer weight computation ↔ train.py, lines 115–182 · score 0.61 · global context vector, global context processing, attention mask, weighted, layer
  6. [6] § STAR★Methods › Method details › The ALF module operates in several stages › Layer weight computation ↔ inference.py, lines 185–229 · score 0.59 · Token scores, token weights, softmaxed, masked, Layer
  7. [7] § STAR★Methods › Method details › The ALF module operates in several stages › Layer weight computation ↔ train.py, lines 253–302 · score 0.59 · Token scores, token weights, softmaxed, masked, Layer
  8. [8] § STAR★Methods › Method details › Model training, inference, and evaluation protocol ↔ train.py, lines 1276–1349 · score 0.59 · LoRA Adapter, random seed, MIMIC IV, patience, reproducibility, epoch
  9. [9] § STAR★Methods › Experimental model and study participant details › Public ICU cohorts and prediction task ↔ train.py, lines 64–80 · score 0.57 · LODS, MELD, OASIS, SIRS, SOFA, admissions
  10. [10] § STAR★Methods › Method details › ALFIA architecture ↔ train.py, lines 342–421 · score 0.55 · Attentional Classifier Head, Adaptive Layer Fusion, LoRA, transformer, module, trained
  11. [11] § Results › Training dynamics demonstrate ALFIA convergence in textual data learning ↔ train.py, lines 1276–1349 · score 0.52 · target modules, validation AUPRC, fusion layers, LoRA, alpha, hyperparameters
  12. [12] § STAR★Methods › Method details › Attentional classifier head (ACH) module ↔ inference.py, lines 231–263 · score 0.51 · Attentional Classifier Head, fused embedding, unsqueezed, module
  13. [13] § STAR★Methods › Method details › ALFIA architecture ↔ inference.py, lines 265–325 · score 0.51 · Attentional Classifier Head, Adaptive Layer Fusion, LoRA, transformer, module, model

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 1,349 lines · 84 KB · Apache-2.0 · 9 matches

  1. # ───── Version and SCRIPT_VERSION ───────────────────────────────────────────
  2. SCRIPT_VERSION = "1.4.3" # Incremented version
  3. # CHANGELOG removed as per user request
  4. import logging
  5. from pathlib import Path
  6. import pandas as pd
  7. import numpy as np
  8. import matplotlib.pyplot as plt
  9. import seaborn as sns
  10. import os
  11. import torch
  12. import torch.nn as nn
  13. import torch.nn.functional as F
  14. import math
  15. from transformers import AutoTokenizer, AutoModel, get_linear_schedule_with_warmup
  16. from torch.optim import AdamW
  17. from torch.utils.data import DataLoader, Dataset
  18. from typing import List, Dict, Tuple, Optional, Union, Any
  19. from tqdm import tqdm
  20. from sklearn.model_selection import train_test_split
  21. from sklearn.metrics import precision_recall_fscore_support, roc_auc_score, confusion_matrix, classification_report, average_precision_score, log_loss
  22. from sklearn.utils import resample # Added for bootstrapping
  23. import json
  24. import random
  25. from datetime import datetime # Added for timestamp
  26. # Import PEFT libraries
  27. try:
  28. from peft import LoraConfig, get_peft_model, TaskType, PeftModel
  29. PEFT_AVAILABLE = True
  30. except ImportError:
  31. PEFT_AVAILABLE = False
  32. logging.warning("PEFT library not found. LoRA functionality will be disabled.")
  33. # ───── logging setup (initial console setup) ───────────────────────────
  34. logging.basicConfig(level=logging.INFO,
  35. format="%(asctime)s | %(levelname)s | %(message)s",
  36. handlers=[logging.StreamHandler()])
  37. logging.info(f"Script Version: {SCRIPT_VERSION}")
  38. # logging.info("Changelog:\n" + CHANGELOG) # Changelog display removed
  39. os.environ['HF_ENDPOINT'] = 'https://hf-mirror.com'
  40. # ───── paths (placeholders, will be redefined in main) ────────────────
  41. INPUT_PATH = Path("data/mimiciv_text.csv")
  42. BASE_RESULTS_DIR = Path("results")
  43. EMBED_PATH: Path = None
  44. FIG_PATH: Path = None
  45. TRAINED_FUSION_PATH: Path = None
  46. TRAINED_LORA_ADAPTER_PATH: Path = None
  47. EPOCH_METRICS_SAVE_PATH: Path = None
  48. LOG_FILE_PATH: Path = None
  49. HYPERPARAMS_FILE_PATH: Path = None
  50. RUN_SPECIFIC_DIR: Path = None
  51. BEST_VAL_PREDS_SAVE_PATH: Path = None
  52. METRICS_CI_VAL_SAVE_PATH: Path = None
  53. BEST_TEST_PREDS_SAVE_PATH: Path = None
  54. METRICS_CI_TEST_SAVE_PATH: Path = None
  55. # ───── Control Parameters ───────────────────────────────────────────
  56. TARGET_COL = "hospital_expire_flag"
  57. TEXT_COL = "patient_description"
  58. EMBED_MODEL = "dmis-lab/biobert-base-cased-v1.2"
  59. GLOBAL_SEED = 42
  60. DROP_COLS = [
  61. "subject_id", "hadm_id", "stay_id", "dischtime", "admittime",
  62. "icu_intime", "icu_outtime", "admission_type", "admission_location",
  63. "discharge_location", "careunit", "los_hospital", "los_icu",
  64. "dod", "mortality_365", "icu_expire_flag", "follow_up_years",
  65. "survival_days", "mortality_28", "mortality_30", "mortality_90",
  66. "apsiii", "sapsii", "sofa", "gcs", "sirs", "lods", "charlson",
  67. "meld", "sepsis_time", "oasis", "InvasiveVent", "cpr", "crrt",
  68. "rrt", "lvef_min", "aki_score", "aki_time", "microbiology", "race", "diagnosis_long_title1"
  69. ]
  70. ALWAYS_KEEP = {"lactate_max"}
  71. # ───── Training Hyperparameters ─────────────────────────
  72. NUM_EPOCHS = 10
  73. TRAIN_BATCH_SIZE = 4
  74. LEARNING_RATE = 0.0001
  75. MAX_SEQ_LENGTH = 512
  76. NUM_FUSION_LAYERS = 4
  77. MIN_SAMPLES_FOR_TOP_LAYER_FINETUNING = 5000
  78. EARLY_STOPPING_PATIENCE = 5
  79. TEST_SPLIT_RATIO_CONST = 0.125
  80. VALIDATION_SPLIT_RATIO_DEFAULT = 0.125 / (1.0 - TEST_SPLIT_RATIO_CONST)
  81. N_BOOTSTRAPS_CI = 1000
  82. # ───── LoRA Hyperparameters (Defaults) ───────────────────
  83. USE_LORA_DEFAULT = False
  84. LORA_R_DEFAULT = 8
  85. LORA_ALPHA_DEFAULT = 16
  86. LORA_DROPOUT_DEFAULT = 0.05
  87. LORA_TARGET_MODULES_DEFAULT = "query,key,value"
  88. # ───── Reproducibility Function ─────────────────────────
  89. def set_seed(seed_value: int = GLOBAL_SEED):
  90. random.seed(seed_value)
  91. np.random.seed(seed_value)
  92. torch.manual_seed(seed_value)
  93. if torch.cuda.is_available():
  94. torch.cuda.manual_seed_all(seed_value)
  95. torch.backends.cudnn.deterministic = True
  96. torch.backends.cudnn.benchmark = False
  97. os.environ['PYTHONHASHSEED'] = str(seed_value)
  98. logging.info(f"Global seed set to {seed_value}")
  99. # ───── Adaptive Layer Fusion Module (Modified - Uses Average Pooling) ──────────────────────────────────
  100. class AdaptiveLayerFusion(nn.Module):
  101. def __init__(
  102. self,
  103. hidden_size: int,
  104. num_layers: int = 4, # Number of layers to fuse from the top of the base model
  105. num_heads: int = 4, # Number of attention heads for layer attention mechanism
  106. dropout: float = 0.1,
  107. use_layer_gating: bool = True # Whether to use a sigmoid gate on layer weights
  108. ):
  109. super().__init__()
  110. self.hidden_size = hidden_size
  111. self.num_heads = num_heads
  112. self.use_layer_gating = use_layer_gating
  113. self.num_layers_to_fuse = num_layers # Stores the intended number of layers to fuse
  114. head_dim = 64 # Fixed dimension for each head in layer attention
  115. self.layer_query_projection = nn.Linear(hidden_size, num_heads * head_dim)
  116. self.layer_key_projection = nn.Linear(hidden_size, num_heads * head_dim)
  117. self.layer_value_projection = nn.Linear(hidden_size, num_heads * head_dim)
  118. self.layer_attention_dropout = nn.Dropout(dropout)
  119. self.layer_attention_output = nn.Linear(num_heads * head_dim, hidden_size)
  120. self.layer_interaction = nn.Sequential(
  121. nn.Linear(hidden_size, hidden_size),
  122. nn.LayerNorm(hidden_size),
  123. nn.GELU(),
  124. nn.Dropout(dropout),
  125. nn.Linear(hidden_size, hidden_size)
  126. )
  127. if use_layer_gating:
  128. self.layer_gate = nn.Sequential(
  129. nn.Linear(hidden_size, self.num_layers_to_fuse), # Output dim matches number of layers to fuse
  130. nn.Sigmoid()
  131. )
  132. self.content_projection = nn.Linear(hidden_size, hidden_size)
  133. self.layer_norm1 = nn.LayerNorm(hidden_size)
  134. # Token-level attention mechanism to get a local context vector
  135. self.token_attention = nn.Sequential(
  136. nn.Linear(hidden_size, hidden_size // 2),
  137. nn.LayerNorm(hidden_size // 2),
  138. nn.GELU(),
  139. nn.Dropout(dropout/2), # Reduced dropout for this part
  140. nn.Linear(hidden_size // 2, hidden_size // 4),
  141. nn.GELU(),
  142. nn.Linear(hidden_size // 4, 1)
  143. )
  144. self.token_attention_dropout = nn.Dropout(dropout)
  145. # Global context processing layer (operates on average pooled representation)
  146. self.global_context_layer = nn.Sequential(
  147. nn.Linear(hidden_size, hidden_size),
  148. nn.LayerNorm(hidden_size),
  149. nn.GELU(),
  150. nn.Dropout(dropout)
  151. )
  152. # Fuses local (token-attended) and global (average-pooled) contexts
  153. self.context_fusion = nn.Sequential(
  154. nn.Linear(hidden_size * 2, hidden_size), # Input is concatenation of local and global
  155. nn.LayerNorm(hidden_size),
  156. nn.Dropout(dropout)
  157. )
  158. self.output_projection = nn.Linear(hidden_size, hidden_size)
  159. self.layer_norm2 = nn.LayerNorm(hidden_size)
  160. self.register_buffer('last_attention_mask', None, persistent=False) # Stores the attention mask from the last forward pass
  161. def _compute_layer_weights(self, hidden_states: List[torch.Tensor]) -> torch.Tensor:
  162. """
  163. Computes attention weights over the input hidden_states (from different layers of a base model).
  164. Uses average-pooled representations of each layer as keys and values for an attention mechanism.
  165. The query is derived from the mean of these average-pooled representations.
  166. """
  167. batch_size = hidden_states[0].size(0)
  168. num_layers_being_fused = len(hidden_states) # Actual number of layers passed in
  169. head_dim = self.layer_query_projection.out_features // self.num_heads
  170. # Use the stored attention_mask to perform masked average pooling for each layer
  171. if self.last_attention_mask is not None:
  172. attention_mask = self.last_attention_mask
  173. mask_expanded = attention_mask.unsqueeze(-1).float() # (batch, seq_len, 1)
  174. avg_pooled_layers = []
  175. for layer in hidden_states:
  176. masked_layer = layer * mask_expanded # Apply mask
  177. seq_lengths = attention_mask.sum(dim=1, keepdim=True).float() # (batch, 1)
  178. avg_pooled = masked_layer.sum(dim=1) / seq_lengths.clamp(min=1) # (batch, hidden_size)
  179. avg_pooled_layers.append(avg_pooled)
  180. avg_reps = torch.stack(avg_pooled_layers, dim=1) # (batch, num_layers_being_fused, hidden_size)
  181. else:
  182. # Fallback if mask is not available (should not happen in normal operation)
  183. logging.warning("last_attention_mask not found in AdaptiveLayerFusion. Using simple mean pooling for layer weights.")
  184. avg_pooled_layers = [layer.mean(dim=1) for layer in hidden_states]
  185. avg_reps = torch.stack(avg_pooled_layers, dim=1)
  186. query_input = avg_reps.mean(dim=1) # (batch, hidden_size), query derived from average of all layer representations
  187. query = self.layer_query_projection(query_input).view(batch_size, self.num_heads, head_dim) # (batch, num_heads, head_dim)
  188. key = self.layer_key_projection(avg_reps).view(batch_size, num_layers_being_fused, self.num_heads, head_dim) # (batch, num_layers, num_heads, head_dim)
  189. value = self.layer_value_projection(avg_reps).view(batch_size, num_layers_being_fused, self.num_heads, head_dim) # (batch, num_layers, num_heads, head_dim)
  190. key = key.permute(0, 2, 1, 3) # (batch, num_heads, num_layers, head_dim)
  191. value = value.permute(0, 2, 1, 3) # (batch, num_heads, num_layers, head_dim)
  192. attention_scores = torch.matmul(query.unsqueeze(2), key.transpose(-1, -2)) / math.sqrt(head_dim) # (batch, num_heads, 1, num_layers)
  193. attention_weights = F.softmax(attention_scores, dim=-1) # (batch, num_heads, 1, num_layers)
  194. attention_weights = self.layer_attention_dropout(attention_weights)
  195. context = torch.matmul(attention_weights, value) # (batch, num_heads, 1, head_dim)
  196. context = context.squeeze(2).contiguous().view(batch_size, self.num_heads * head_dim) # (batch, num_heads * head_dim)
  197. context = self.layer_attention_output(context) # (batch, hidden_size)
  198. if self.use_layer_gating:
  199. # Dynamically adjust layer_gate's output dimension if it doesn't match num_layers_being_fused
  200. if self.layer_gate[0].out_features != num_layers_being_fused:
  201. logging.warning(
  202. f"AdaptiveLayerFusion's layer_gate was initialized for {self.layer_gate[0].out_features} layers, "
  203. f"but received {num_layers_being_fused} layers in forward pass. "
  204. f"Re-initializing layer_gate's final linear layer for {num_layers_being_fused} outputs."
  205. )
  206. original_device = self.layer_gate[0].weight.device
  207. original_dtype = self.layer_gate[0].weight.dtype
  208. self.layer_gate[0] = nn.Linear(self.hidden_size, num_layers_being_fused).to(device=original_device, dtype=original_dtype)
  209. gate_weights = self.layer_gate(context) # (batch, num_layers_being_fused)
  210. # Sanity check after potential dynamic adjustment
  211. if gate_weights.shape[1] != num_layers_being_fused:
  212. logging.error(f"CRITICAL: Layer gate output features ({gate_weights.shape[1]}) "
  213. f"still does not match number of layers being fused ({num_layers_being_fused}) "
  214. "after attempted dynamic adjustment. Falling back to uniform weights.")
  215. return torch.ones(batch_size, num_layers_being_fused, device=context.device) / num_layers_being_fused
  216. return gate_weights # (batch_size, num_layers_being_fused)
  217. else:
  218. # If not using gating, return the averaged attention weights across heads
  219. return attention_weights.squeeze(2).mean(dim=1) # (batch_size, num_layers_being_fused)
  220. def forward(
  221. self,
  222. all_hidden_states: List[torch.Tensor], # List of (batch, seq_len, hidden_size) tensors
  223. attention_mask: torch.Tensor # (batch, seq_len)
  224. ) -> Tuple[torch.Tensor, torch.Tensor]: # (pooled_output, layer_weights)
  225. if not all_hidden_states:
  226. raise ValueError("all_hidden_states list cannot be empty.")
  227. self.last_attention_mask = attention_mask # Store for _compute_layer_weights
  228. layer_weights = self._compute_layer_weights(all_hidden_states) # (batch, num_fused_layers)
  229. # Fallback if layer_weights shape is incorrect (e.g., due to dynamic adjustment failure)
  230. if layer_weights.shape[1] != len(all_hidden_states):
  231. logging.error(f"Shape mismatch after _compute_layer_weights: layer_weights has {layer_weights.shape[1]} weights, "
  232. f"but {len(all_hidden_states)} hidden states were provided. Using uniform weights as fallback.")
  233. num_fused_layers = len(all_hidden_states)
  234. layer_weights = torch.ones(all_hidden_states[0].size(0), num_fused_layers, device=all_hidden_states[0].device) / num_fused_layers
  235. # Weighted sum of layer hidden states
  236. weighted_states = torch.zeros_like(all_hidden_states[0]) # (batch, seq_len, hidden_size)
  237. for i, layer_state in enumerate(all_hidden_states):
  238. current_layer_weight = layer_weights[:, i].unsqueeze(1).unsqueeze(2) # (batch, 1, 1)
  239. weighted_states += current_layer_weight * layer_state
  240. interaction_states = self.layer_interaction(weighted_states) # (batch, seq_len, hidden_size)
  241. projected_states = self.content_projection(weighted_states)
  242. enhanced_states = self.layer_norm1(weighted_states + projected_states + interaction_states) # (batch, seq_len, hidden_size)
  243. # Token-level attention for local context
  244. token_scores = self.token_attention(enhanced_states).squeeze(-1) # (batch, seq_len)
  245. token_scores = token_scores.masked_fill(attention_mask == 0, -1e9) # Apply attention mask
  246. token_weights = F.softmax(token_scores, dim=1) # (batch, seq_len)
  247. token_weights = self.token_attention_dropout(token_weights)
  248. local_context = torch.bmm(token_weights.unsqueeze(1), enhanced_states).squeeze(1) # (batch, hidden_size)
  249. # Global context (masked average pooling over enhanced_states)
  250. mask_expanded = attention_mask.unsqueeze(-1).float()
  251. masked_enhanced_states = enhanced_states * mask_expanded
  252. seq_lengths = attention_mask.sum(dim=1, keepdim=True).float().clamp(min=1)
  253. global_context_unprocessed = masked_enhanced_states.sum(dim=1) / seq_lengths # (batch, hidden_size)
  254. global_context = self.global_context_layer(global_context_unprocessed) # (batch, hidden_size)
  255. # Combine local and global contexts
  256. combined_context = torch.cat([local_context, global_context], dim=1) # (batch, hidden_size * 2)
  257. pooled_output = self.context_fusion(combined_context) # (batch, hidden_size)
  258. projected_output = self.output_projection(pooled_output)
  259. pooled_output = self.layer_norm2(pooled_output + projected_output) # Final pooled output
  260. return pooled_output, layer_weights
  261. # ───── Attentional Classifier Head ───────────────────────────────────
  262. class AttentionalClassifierHead(nn.Module):
  263. def __init__(self, hidden_size: int, num_classes: int, num_attn_heads: int = 4, ffn_expansion: int = 2, dropout_rate: float = 0.1):
  264. super().__init__()
  265. # Adjust num_attn_heads if hidden_size is not divisible by it
  266. if hidden_size % num_attn_heads != 0:
  267. valid_heads = [h for h in [16, 12, 8, 4, 2, 1] if hidden_size % h == 0] # Common head counts
  268. if not valid_heads: # Should not happen if hidden_size is typical (e.g., 768)
  269. raise ValueError(f"The hidden size ({hidden_size}) is not divisible by any common attention head counts (1,2,4,8,12,16).")
  270. original_num_attn_heads = num_attn_heads
  271. num_attn_heads = valid_heads[0] # Pick the largest valid head count
  272. logging.warning(f"AttentionalClassifierHead: num_attn_heads ({original_num_attn_heads}) caused hidden_size ({hidden_size}) not to be divisible. Adjusted to {num_attn_heads}.")
  273. self.hidden_size = hidden_size
  274. self.num_attn_heads = num_attn_heads
  275. self.attention = nn.MultiheadAttention(embed_dim=hidden_size, num_heads=self.num_attn_heads, dropout=dropout_rate, batch_first=True)
  276. self.norm1 = nn.LayerNorm(hidden_size)
  277. self.norm2 = nn.LayerNorm(hidden_size)
  278. self.ffn = nn.Sequential(
  279. nn.Linear(hidden_size, hidden_size * ffn_expansion),
  280. nn.GELU(),
  281. nn.Dropout(dropout_rate),
  282. nn.Linear(hidden_size * ffn_expansion, hidden_size)
  283. )
  284. self.dropout = nn.Dropout(dropout_rate)
  285. self.output_layer = nn.Linear(hidden_size, num_classes)
  286. def forward(self, fused_embeddings: torch.Tensor) -> torch.Tensor: # fused_embeddings: (batch, hidden_size)
  287. x = fused_embeddings.unsqueeze(1) # (batch, 1, hidden_size) - MHA expects sequence
  288. attn_output, _ = self.attention(query=x, key=x, value=x) # Self-attention on the single fused vector
  289. x = self.norm1(x + self.dropout(attn_output))
  290. ffn_output = self.ffn(x)
  291. x = self.norm2(x + self.dropout(ffn_output))
  292. x_squeezed = x.squeeze(1) # (batch, hidden_size)
  293. return self.output_layer(x_squeezed)
  294. # ───── Complete Classification Model ─────────────────────────────────────
  295. class TextEmbedderWithClassifier(nn.Module):
  296. def __init__(self,
  297. base_model_name: str,
  298. num_fusion_layers: int, # How many top layers from base_model to fuse
  299. fusion_dropout: float = 0.1,
  300. num_classes: int = 1, # For binary classification with BCEWithLogitsLoss
  301. use_lora: bool = False,
  302. lora_config_dict: Optional[Dict[str, Any]] = None,
  303. classifier_num_attn_heads: int = 4, # For AttentionalClassifierHead
  304. classifier_ffn_expansion: int = 2): # For AttentionalClassifierHead
  305. super().__init__()
  306. self.num_fusion_layers = num_fusion_layers
  307. self.use_lora = use_lora and PEFT_AVAILABLE
  308. self.base_model_internal = AutoModel.from_pretrained(base_model_name, output_hidden_states=True)
  309. self.num_base_model_total_hidden_states = self.base_model_internal.config.num_hidden_layers + 1 # Includes embedding layer
  310. if self.use_lora:
  311. if lora_config_dict is None: lora_config_dict = {}
  312. raw_target_modules = lora_config_dict.get("target_modules", LORA_TARGET_MODULES_DEFAULT)
  313. if isinstance(raw_target_modules, str):
  314. target_modules_list = [m.strip() for m in raw_target_modules.split(',') if m.strip()]
  315. elif isinstance(raw_target_modules, list):
  316. target_modules_list = raw_target_modules
  317. else: # Fallback to default if type is unexpected
  318. target_modules_list = [m.strip() for m in LORA_TARGET_MODULES_DEFAULT.split(',') if m.strip()]
  319. lora_config = LoraConfig(
  320. r=lora_config_dict.get("r", LORA_R_DEFAULT),
  321. lora_alpha=lora_config_dict.get("lora_alpha", LORA_ALPHA_DEFAULT),
  322. target_modules=target_modules_list,
  323. lora_dropout=lora_config_dict.get("lora_dropout", LORA_DROPOUT_DEFAULT),
  324. bias="none", # Common LoRA setting
  325. task_type=TaskType.FEATURE_EXTRACTION # Or SEQ_CLS if LoRA is applied to a model with a classification head
  326. )
  327. self.base_model = get_peft_model(self.base_model_internal, lora_config)
  328. logging.info("LoRA applied to the base model.")
  329. self.base_model.print_trainable_parameters()
  330. else:
  331. self.base_model = self.base_model_internal
  332. if use_lora and not PEFT_AVAILABLE: # Log if LoRA was intended but not possible
  333. logging.warning("LoRA requested but PEFT library not available. Proceeding without LoRA.")
  334. self.fusion_module = AdaptiveLayerFusion(
  335. hidden_size=self.base_model.config.hidden_size,
  336. num_layers=num_fusion_layers, # This is the `num_layers_to_fuse` for the gate initialization
  337. num_heads=4, # Default from AdaptiveLayerFusion
  338. dropout=fusion_dropout
  339. )
  340. self.classifier = AttentionalClassifierHead(
  341. hidden_size=self.base_model.config.hidden_size, # Input to classifier is output of fusion
  342. num_classes=num_classes,
  343. num_attn_heads=classifier_num_attn_heads,
  344. ffn_expansion=classifier_ffn_expansion,
  345. dropout_rate=fusion_dropout # Use same dropout as fusion for consistency
  346. )
  347. logging.info(f"TextEmbedder: Using AttentionalClassifierHead with num_attn_heads={self.classifier.num_attn_heads}, ffn_expansion={classifier_ffn_expansion}, dropout={fusion_dropout}")
  348. def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor):
  349. base_model_outputs = self.base_model(input_ids=input_ids, attention_mask=attention_mask)
  350. all_layers_hidden_states = list(base_model_outputs.hidden_states) # Tuple to list
  351. # Validate num_fusion_layers against actual available transformer layers
  352. actual_num_transformer_layers = len(all_layers_hidden_states) -1 # Exclude initial embedding layer output
  353. if not (0 < self.num_fusion_layers <= actual_num_transformer_layers):
  354. raise ValueError(
  355. f"num_fusion_layers ({self.num_fusion_layers}) is invalid. "
  356. f"The base model has {actual_num_transformer_layers} transformer layers. "
  357. f"num_fusion_layers must be > 0 and <= {actual_num_transformer_layers}."
  358. )
  359. # Select the top `num_fusion_layers` from the transformer layers (excluding embedding layer)
  360. # `all_layers_hidden_states` includes embedding layer as [0], then transformer layers [1]...[N]
  361. # So, to get the top K transformer layers, we take from the end.
  362. hidden_states_for_fusion = all_layers_hidden_states[-self.num_fusion_layers:]
  363. fused_embeddings, _ = self.fusion_module(hidden_states_for_fusion, attention_mask)
  364. return self.classifier(fused_embeddings)
  365. # ───── Custom Dataset ───────────────────────────────────────
  366. class MIMICIVTextDataset(Dataset):
  367. def __init__(self, texts: List[str], labels: List[int], tokenizer, max_len: int):
  368. self.texts = texts
  369. self.labels = labels
  370. self.tokenizer = tokenizer
  371. self.max_len = max_len
  372. def __len__(self):
  373. return len(self.texts)
  374. def __getitem__(self, item_idx):
  375. text = str(self.texts[item_idx])
  376. label = self.labels[item_idx]
  377. encoding = self.tokenizer.encode_plus(
  378. text,
  379. add_special_tokens=True,
  380. max_length=self.max_len,
  381. return_token_type_ids=False, # Not needed for most BERT-like models
  382. padding='max_length', # Pad to max_len
  383. truncation=True, # Truncate to max_len
  384. return_attention_mask=True,
  385. return_tensors='pt', # Return PyTorch tensors
  386. )
  387. return {
  388. 'text': text, # Keep original text for reference if needed
  389. 'input_ids': encoding['input_ids'].flatten(),
  390. 'attention_mask': encoding['attention_mask'].flatten(),
  391. 'labels': torch.tensor(label, dtype=torch.float) # For BCEWithLogitsLoss
  392. }
  393. # ───── Compute Evaluation Metrics Function ─────────────────────────────────
  394. def compute_metrics(preds: torch.Tensor, targets: torch.Tensor, threshold: float = 0.5) -> Dict[str, float]:
  395. preds_cpu = preds.detach().cpu().numpy()
  396. targets_cpu = targets.detach().cpu().numpy()
  397. # Ensure preds_cpu is 1D for binary classification probability scores
  398. if preds_cpu.ndim > 1 and preds_cpu.shape[1] == 1:
  399. preds_cpu = preds_cpu.flatten()
  400. binary_preds_np = (preds_cpu >= threshold).astype(int)
  401. accuracy = (binary_preds_np == targets_cpu).mean().item()
  402. precision, recall, f1, _ = precision_recall_fscore_support(
  403. targets_cpu, binary_preds_np, average='binary', zero_division=0
  404. )
  405. cm = confusion_matrix(targets_cpu, binary_preds_np, labels=[0, 1]) # Ensure labels are [0,1]
  406. # Handle cases where confusion matrix might not be 2x2 (e.g., all preds/targets are same class)
  407. if cm.size == 4: # Standard 2x2 case
  408. tn, fp, fn, tp = cm.ravel()
  409. else: # Degenerate cases
  410. tn, fp, fn, tp = 0, 0, 0, 0
  411. if (targets_cpu == 0).all() and (binary_preds_np == 0).all(): tn = len(targets_cpu)
  412. elif (targets_cpu == 1).all() and (binary_preds_np == 1).all(): tp = len(targets_cpu)
  413. elif (targets_cpu == 0).all() and (binary_preds_np == 1).all(): fp = len(targets_cpu) # All actual negatives predicted positive
  414. elif (targets_cpu == 1).all() and (binary_preds_np == 0).all(): fn = len(targets_cpu) # All actual positives predicted negative
  415. # Other combinations (e.g. mixed actuals, all same predictions) are also possible but covered by initializing to 0
  416. specificity = tn / (tn + fp) if (tn + fp) > 0 else 0.0
  417. auc_roc = float('nan')
  418. auprc = float('nan')
  419. # AUC-ROC and AUPRC require at least two classes in targets
  420. if len(np.unique(targets_cpu)) > 1:
  421. try:
  422. auc_roc = roc_auc_score(targets_cpu, preds_cpu) # Use probabilities for AUCs
  423. except ValueError as e:
  424. logging.debug(f"AUC-ROC calculation error: {e}. Targets unique: {np.unique(targets_cpu)}, Preds shape: {preds_cpu.shape}")
  425. try:
  426. auprc = average_precision_score(targets_cpu, preds_cpu) # Use probabilities for AUCs
  427. except ValueError as e:
  428. logging.debug(f"AUPRC calculation error: {e}. Targets unique: {np.unique(targets_cpu)}, Preds shape: {preds_cpu.shape}")
  429. else:
  430. logging.debug("Single class in targets, AUC-ROC and AUPRC are not defined (NaN).")
  431. return {
  432. 'accuracy': accuracy,
  433. 'precision': precision,
  434. 'recall': recall, # Sensitivity
  435. 'f1': f1,
  436. 'specificity': specificity,
  437. 'balanced_acc': (recall + specificity) / 2.0 if recall is not None and specificity is not None else 0.0,
  438. 'auc_roc': auc_roc,
  439. 'auprc': auprc,
  440. 'ppv': tp / (tp + fp) if (tp + fp) > 0 else 0.0, # Positive Predictive Value (same as precision)
  441. 'npv': tn / (tn + fn) if (tn + fn) > 0 else 0.0, # Negative Predictive Value
  442. 'pos_ratio_actual': targets_cpu.mean().item(), # Actual positive rate
  443. 'neg_ratio_actual': 1.0 - targets_cpu.mean().item(), # Actual negative rate
  444. 'tp': int(tp), 'fp': int(fp), 'tn': int(tn), 'fn': int(fn)
  445. }
  446. # ───── Calculate Evaluation Metrics with 95% CI Function ─────────────────────────
  447. def calculate_metrics_with_ci(
  448. y_true: np.ndarray,
  449. y_pred_proba: np.ndarray,
  450. threshold: float = 0.5,
  451. n_bootstraps: int = N_BOOTSTRAPS_CI, # Global default
  452. alpha: float = 0.05 # For 95% CI
  453. ) -> pd.DataFrame:
  454. y_true = np.asarray(y_true)
  455. y_pred_proba = np.asarray(y_pred_proba)
  456. # Ensure y_pred_proba is 1D
  457. if y_pred_proba.ndim > 1 and y_pred_proba.shape[1] == 1:
  458. y_pred_proba = y_pred_proba.flatten()
  459. metrics_results_list = []
  460. bootstrapped_metrics_values = {
  461. 'roc_auc': [], 'prauc': [], 'logloss': [],
  462. 'accuracy': [], 'f1': [], 'precision': [], 'recall': [], 'specificity': []
  463. }
  464. eps = 1e-15 # For clipping probabilities in log_loss
  465. y_pred_proba_clipped = np.clip(y_pred_proba, eps, 1 - eps)
  466. # Calculate point estimates first
  467. point_estimates = {}
  468. y_pred_binary_orig = (y_pred_proba >= threshold).astype(int)
  469. if len(np.unique(y_true)) > 1:
  470. try: point_estimates['roc_auc'] = roc_auc_score(y_true, y_pred_proba)
  471. except ValueError: point_estimates['roc_auc'] = np.nan
  472. try: point_estimates['prauc'] = average_precision_score(y_true, y_pred_proba)
  473. except ValueError: point_estimates['prauc'] = np.nan
  474. else:
  475. point_estimates['roc_auc'] = np.nan
  476. point_estimates['prauc'] = np.nan
  477. try: point_estimates['logloss'] = log_loss(y_true, y_pred_proba_clipped, labels=[0,1]) # Ensure labels are specified for consistency
  478. except ValueError: point_estimates['logloss'] = np.nan
  479. precision_orig, recall_orig, f1_orig, _ = precision_recall_fscore_support(
  480. y_true, y_pred_binary_orig, average='binary', zero_division=0
  481. )
  482. cm_orig = confusion_matrix(y_true, y_pred_binary_orig, labels=[0,1])
  483. if cm_orig.size == 4: tn_orig, fp_orig, fn_orig, tp_orig = cm_orig.ravel()
  484. else: tn_orig, fp_orig, fn_orig, tp_orig = 0,0,0,0 # Handle degenerate cases
  485. point_estimates['accuracy'] = (y_pred_binary_orig == y_true).mean()
  486. point_estimates['f1'] = f1_orig
  487. point_estimates['precision'] = precision_orig
  488. point_estimates['recall'] = recall_orig
  489. point_estimates['specificity'] = tn_orig / (tn_orig + fp_orig) if (tn_orig + fp_orig) > 0 else 0.0
  490. # Bootstrapping
  491. rng = np.random.RandomState(GLOBAL_SEED) # Use global seed for reproducibility of bootstrapping
  492. for _ in tqdm(range(n_bootstraps), desc="Bootstrapping CI", leave=False):
  493. indices = rng.choice(len(y_true), size=len(y_true), replace=True)
  494. if len(indices) == 0: continue # Should not happen if y_true is not empty
  495. y_true_boot = y_true[indices]
  496. y_pred_proba_boot = y_pred_proba[indices]
  497. y_pred_proba_boot_clipped = np.clip(y_pred_proba_boot, eps, 1 - eps)
  498. y_pred_binary_boot = (y_pred_proba_boot >= threshold).astype(int)
  499. if len(np.unique(y_true_boot)) > 1:
  500. try: bootstrapped_metrics_values['roc_auc'].append(roc_auc_score(y_true_boot, y_pred_proba_boot))
  501. except ValueError: bootstrapped_metrics_values['roc_auc'].append(np.nan)
  502. try: bootstrapped_metrics_values['prauc'].append(average_precision_score(y_true_boot, y_pred_proba_boot))
  503. except ValueError: bootstrapped_metrics_values['prauc'].append(np.nan)
  504. else:
  505. bootstrapped_metrics_values['roc_auc'].append(np.nan)
  506. bootstrapped_metrics_values['prauc'].append(np.nan)
  507. try: bootstrapped_metrics_values['logloss'].append(log_loss(y_true_boot, y_pred_proba_boot_clipped, labels=[0,1]))
  508. except ValueError: bootstrapped_metrics_values['logloss'].append(np.nan)
  509. precision_boot, recall_boot, f1_boot, _ = precision_recall_fscore_support(
  510. y_true_boot, y_pred_binary_boot, average='binary', zero_division=0
  511. )
  512. cm_boot = confusion_matrix(y_true_boot, y_pred_binary_boot, labels=[0,1])
  513. if cm_boot.size == 4: tn_boot, fp_boot, fn_boot, tp_boot = cm_boot.ravel()
  514. else: tn_boot, fp_boot, fn_boot, tp_boot = 0,0,0,0
  515. bootstrapped_metrics_values['accuracy'].append((y_pred_binary_boot == y_true_boot).mean())
  516. bootstrapped_metrics_values['f1'].append(f1_boot)
  517. bootstrapped_metrics_values['precision'].append(precision_boot)
  518. bootstrapped_metrics_values['recall'].append(recall_boot)
  519. bootstrapped_metrics_values['specificity'].append(tn_boot / (tn_boot + fp_boot) if (tn_boot + fp_boot) > 0 else 0.0)
  520. # Calculate CIs from bootstrapped values
  521. for metric_name, point_estimate_val in point_estimates.items():
  522. boot_values_for_metric = [v for v in bootstrapped_metrics_values[metric_name] if not np.isnan(v)] # Filter out NaNs
  523. ci_lower, ci_upper = np.nan, np.nan
  524. if len(boot_values_for_metric) > 1: # Need at least 2 values for percentile
  525. ci_lower = np.percentile(boot_values_for_metric, (alpha / 2) * 100)
  526. ci_upper = np.percentile(boot_values_for_metric, (1 - alpha / 2) * 100)
  527. metrics_results_list.append({
  528. 'metric': metric_name,
  529. 'value': point_estimate_val,
  530. 'ci_lower (95%)': ci_lower,
  531. 'ci_upper (95%)': ci_upper
  532. })
  533. return pd.DataFrame(metrics_results_list)
  534. # ───── DataLoader worker_init_fn for reproducibility ─────────────────
  535. def seed_worker(worker_id):
  536. # Ensures that each worker in DataLoader has a different, but reproducible, seed
  537. worker_seed = torch.initial_seed() % 2**32
  538. np.random.seed(worker_seed)
  539. random.seed(worker_seed)
  540. # ───── Training Function ──────────────────────────────────
  541. def train_and_save_fusion_model(
  542. texts: List[str], labels: List[int], base_model_name: str,
  543. num_fusion_layers_to_train: int, fusion_model_save_path: Path,
  544. lora_adapter_save_path: Path, use_lora: bool, lora_config_dict: Dict[str, Any],
  545. validation_split_ratio: float, # Proportion of (train+val) data to use for validation
  546. tokenizer_max_len: int = MAX_SEQ_LENGTH,
  547. train_batch_size: int = TRAIN_BATCH_SIZE, num_epochs: int = NUM_EPOCHS,
  548. learning_rate: float = LEARNING_RATE, early_stopping_patience: int = EARLY_STOPPING_PATIENCE,
  549. seed: int = GLOBAL_SEED,
  550. test_dataloader_main: Optional[DataLoader] = None # Kept for structure, but test_dataloader is now created inside
  551. ):
  552. device = "cuda" if torch.cuda.is_available() else "cpu"
  553. logging.info(f"Training on {device} for base model: {base_model_name}, fusing top {num_fusion_layers_to_train} layers.")
  554. if use_lora and PEFT_AVAILABLE: logging.info(f"LoRA ENABLED. Config: {lora_config_dict}. Adapters to: {lora_adapter_save_path.parent}")
  555. else: logging.info("LoRA DISABLED.")
  556. tokenizer = AutoTokenizer.from_pretrained(base_model_name)
  557. texts_np_all, labels_np_all = np.array(texts), np.array(labels)
  558. set_seed(seed) # Set seed before any data splitting or model initialization
  559. # Split into (Train+Val) and Test first
  560. stratify_main_split = labels_np_all if len(np.unique(labels_np_all)) > 1 else None
  561. train_val_texts_np, test_texts_np, train_val_labels_np, test_labels_np = train_test_split(
  562. texts_np_all, labels_np_all, test_size=TEST_SPLIT_RATIO_CONST, random_state=seed, stratify=stratify_main_split
  563. )
  564. train_texts_list: List[str]
  565. train_labels_list: List[int]
  566. val_texts_list: List[str] = []
  567. val_labels_list: List[int] = []
  568. # Split (Train+Val) into Train and Val if validation_split_ratio > 0
  569. if validation_split_ratio > 0 and len(train_val_texts_np) > 0 :
  570. stratify_val_split = train_val_labels_np if len(np.unique(train_val_labels_np)) > 1 else None
  571. train_texts_np_final, val_texts_np_final, train_labels_np_final, val_labels_np_final = train_test_split(
  572. train_val_texts_np, train_val_labels_np, test_size=validation_split_ratio, random_state=seed, stratify=stratify_val_split
  573. )
  574. train_texts_list = train_texts_np_final.tolist()
  575. train_labels_list = train_labels_np_final.tolist()
  576. val_texts_list = val_texts_np_final.tolist()
  577. val_labels_list = val_labels_np_final.tolist()
  578. else: # No validation split, all train_val data goes to training
  579. train_texts_list = train_val_texts_np.tolist()
  580. train_labels_list = train_val_labels_np.tolist()
  581. test_texts_list = test_texts_np.tolist()
  582. test_labels_list = test_labels_np.tolist()
  583. logging.info(f"Data split: {len(train_texts_list)} Train, {len(val_texts_list)} Val, {len(test_texts_list)} Test.")
  584. train_dataset = MIMICIVTextDataset(train_texts_list, train_labels_list, tokenizer, tokenizer_max_len)
  585. g = torch.Generator(); g.manual_seed(seed) # Generator for DataLoader shuffling reproducibility
  586. train_dataloader = DataLoader(train_dataset, batch_size=train_batch_size, shuffle=True,
  587. num_workers=min(2, os.cpu_count()//2 if os.cpu_count() else 1), # Sensible num_workers
  588. pin_memory=(device=="cuda"), worker_init_fn=seed_worker, generator=g)
  589. val_dataloader = None
  590. if val_texts_list:
  591. val_dataset = MIMICIVTextDataset(val_texts_list, val_labels_list, tokenizer, tokenizer_max_len)
  592. val_dataloader = DataLoader(val_dataset, batch_size=train_batch_size, # Use train_batch_size for val for consistency
  593. num_workers=min(2, os.cpu_count()//2 if os.cpu_count() else 1),
  594. pin_memory=(device=="cuda"), worker_init_fn=seed_worker) # No shuffle for val
  595. test_dataloader = None # Initialize test_dataloader
  596. if test_texts_list:
  597. test_dataset = MIMICIVTextDataset(test_texts_list, test_labels_list, tokenizer, tokenizer_max_len)
  598. test_dataloader = DataLoader(test_dataset, batch_size=train_batch_size,
  599. num_workers=min(2, os.cpu_count()//2 if os.cpu_count() else 1),
  600. pin_memory=(device=="cuda"), worker_init_fn=seed_worker) # No shuffle for test
  601. # Calculate positive class weight for imbalanced datasets for BCEWithLogitsLoss
  602. num_pos_train = np.sum(train_labels_list); num_neg_train = len(train_labels_list) - num_pos_train
  603. pos_weight_train = num_neg_train / num_pos_train if num_pos_train > 0 and num_neg_train > 0 else 1.0
  604. logging.info(f"Train Stats: Samples={len(train_labels_list)}, Pos={num_pos_train} ({num_pos_train/len(train_labels_list)*100:.2f}% if len(train_labels_list) > 0 else 0), Neg={num_neg_train}. Pos Weight: {pos_weight_train:.4f}")
  605. if val_dataloader: logging.info(f"Val Stats: Samples={len(val_labels_list)}, Pos={np.sum(val_labels_list)} ({np.sum(val_labels_list)/len(val_labels_list)*100:.2f}% if len(val_labels_list) > 0 else 0)")
  606. if test_dataloader: logging.info(f"Test Stats: Samples={len(test_labels_list)}, Pos={np.sum(test_labels_list)} ({np.sum(test_labels_list)/len(test_labels_list)*100:.2f}% if len(test_labels_list) > 0 else 0)")
  607. model = TextEmbedderWithClassifier(base_model_name, num_fusion_layers_to_train, use_lora=use_lora, lora_config_dict=lora_config_dict).to(device)
  608. # Parameter freezing logic if not using LoRA
  609. if not use_lora or not PEFT_AVAILABLE:
  610. dataset_size = len(texts) # Using the original full dataset size for this decision
  611. if dataset_size >= MIN_SAMPLES_FOR_TOP_LAYER_FINETUNING:
  612. logging.info(f"(No LoRA) Fine-tuning top layers of base model ({dataset_size} samples).")
  613. if hasattr(model.base_model, 'encoder') and hasattr(model.base_model.encoder, 'layer'):
  614. num_actual_layers = len(model.base_model.encoder.layer)
  615. # Unfreeze top 2 transformer layers and pooler (if exists)
  616. layers_to_unfreeze_prefixes = [f"encoder.layer.{num_actual_layers-1}.", f"encoder.layer.{num_actual_layers-2}."]
  617. if hasattr(model.base_model, 'pooler'): layers_to_unfreeze_prefixes.append("pooler.")
  618. for name, param in model.base_model.named_parameters():
  619. param.requires_grad = any(name.startswith(p_prefix) for p_prefix in layers_to_unfreeze_prefixes)
  620. if param.requires_grad: logging.debug(f"Unfreezing (No LoRA): {name}")
  621. else:
  622. logging.warning("(No LoRA) Could not identify encoder layers for partial unfreezing. Base model might be fully frozen or fully trainable.")
  623. else:
  624. logging.info(f"(No LoRA) Freezing entire base model due to small dataset size ({dataset_size} samples). Only fusion/classifier trainable.")
  625. for param in model.base_model.parameters():
  626. param.requires_grad = False
  627. optimizer = AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr=learning_rate, weight_decay=0.01)
  628. scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=int(len(train_dataloader) * num_epochs * 0.1), num_training_steps=len(train_dataloader) * num_epochs)
  629. criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor([pos_weight_train], device=device))
  630. best_val_auprc = 0.0
  631. # These store the state of the best model found so far *in memory* during training.
  632. # Model weights (.pth, LoRA) are saved/overwritten immediately when a new best is found.
  633. # CSVs for metrics/preds are also saved immediately.
  634. best_model_state_in_memory: Optional[Dict[str, Any]] = None # For fusion/classifier weights
  635. best_lora_model_peft_instance_in_memory: Optional[PeftModel] = None # For LoRA adapters
  636. epochs_no_improve = 0
  637. all_epoch_metrics_run = []
  638. logging.info("Starting training loop...")
  639. for epoch in range(num_epochs):
  640. model.train(); total_train_loss = 0
  641. progress_bar = tqdm(train_dataloader, desc=f"Epoch {epoch+1}/{num_epochs} [Train]", leave=False)
  642. for batch in progress_bar:
  643. ids, mask, tgts = batch['input_ids'].to(device), batch['attention_mask'].to(device), batch['labels'].to(device).unsqueeze(1)
  644. optimizer.zero_grad()
  645. outputs = model(input_ids=ids, attention_mask=mask)
  646. loss = criterion(outputs, tgts)
  647. loss.backward()
  648. torch.nn.utils.clip_grad_norm_(filter(lambda p: p.requires_grad, model.parameters()), 1.0) # Gradient clipping
  649. optimizer.step(); scheduler.step()
  650. total_train_loss += loss.item()
  651. progress_bar.set_postfix({'loss': f'{loss.item():.4f}', 'lr': f'{scheduler.get_last_lr()[0]:.2e}'})
  652. avg_train_loss = total_train_loss / len(train_dataloader) if len(train_dataloader) > 0 else 0
  653. logging.info(f"Epoch {epoch+1} - Avg Train Loss: {avg_train_loss:.4f}")
  654. current_epoch_metrics = {'epoch': epoch + 1, 'avg_train_loss': avg_train_loss}
  655. if val_dataloader:
  656. model.eval(); total_val_loss = 0; all_val_preds_epoch_tensors, all_val_labels_epoch_tensors = [], []
  657. val_progress_bar = tqdm(val_dataloader, desc=f"Epoch {epoch+1} [Val]", leave=False)
  658. with torch.no_grad():
  659. for batch in val_progress_bar:
  660. ids, mask, tgts = batch['input_ids'].to(device), batch['attention_mask'].to(device), batch['labels'].to(device).unsqueeze(1)
  661. outputs = model(input_ids=ids, attention_mask=mask)
  662. loss = criterion(outputs, tgts) # Use same criterion for val loss
  663. total_val_loss += loss.item()
  664. all_val_preds_epoch_tensors.append(torch.sigmoid(outputs.detach())) # Store probabilities
  665. all_val_labels_epoch_tensors.append(tgts.detach())
  666. val_progress_bar.set_postfix({'val_loss': f'{loss.item():.4f}'})
  667. avg_val_loss = total_val_loss / len(val_dataloader) if len(val_dataloader) > 0 else 0
  668. current_val_preds_tensor = torch.cat(all_val_preds_epoch_tensors)
  669. current_val_labels_tensor = torch.cat(all_val_labels_epoch_tensors)
  670. val_metrics_epoch = compute_metrics(current_val_preds_tensor, current_val_labels_tensor)
  671. current_epoch_metrics.update({'avg_val_loss': avg_val_loss, **{f'val_{k}': v for k,v in val_metrics_epoch.items()}})
  672. auprc_val = val_metrics_epoch.get('auprc', 0.0); auprc_val = 0.0 if np.isnan(auprc_val) else auprc_val # Handle NaN AUPRC
  673. logging.info(f"Epoch {epoch+1} - Val Loss: {avg_val_loss:.4f}, Val F1: {val_metrics_epoch.get('f1',0):.4f}, Val AUC: {val_metrics_epoch.get('auc_roc',0):.4f}, Val AUPRC: {auprc_val:.4f}")
  674. if auprc_val > best_val_auprc:
  675. best_val_auprc = auprc_val
  676. # Store the model state IN MEMORY first (for potential end-of-training save if logic changes)
  677. best_model_state_in_memory = {
  678. 'fusion_module_state_dict': model.fusion_module.state_dict(),
  679. 'classifier_head_state_dict': model.classifier.state_dict()
  680. }
  681. if use_lora and PEFT_AVAILABLE:
  682. best_lora_model_peft_instance_in_memory = model.base_model # This is a PeftModel instance
  683. epochs_no_improve = 0
  684. logging.info(f"*** New best Val AUPRC: {best_val_auprc:.4f} at Epoch {epoch+1}. Saving weights and results... ***")
  685. # --- Immediately Save/Overwrite Model Weights for this new best model ---
  686. current_best_model_weights_to_save = {
  687. 'fusion_module_state_dict': model.fusion_module.state_dict(),
  688. 'classifier_head_state_dict': model.classifier.state_dict()
  689. }
  690. torch.save(current_best_model_weights_to_save, fusion_model_save_path) # Overwrites
  691. logging.info(f"Best fusion/classifier weights (Epoch {epoch+1}) saved to {fusion_model_save_path}")
  692. if use_lora and PEFT_AVAILABLE:
  693. model.base_model.save_pretrained(str(lora_adapter_save_path)) # Overwrites
  694. logging.info(f"Best LoRA adapters (Epoch {epoch+1}) saved to {lora_adapter_save_path}")
  695. # --- Save/Overwrite Validation Results for this new best model ---
  696. current_best_val_preds_np = current_val_preds_tensor.cpu().numpy().flatten()
  697. current_best_val_labels_np = current_val_labels_tensor.cpu().numpy().flatten()
  698. val_preds_df = pd.DataFrame({'true_labels': current_best_val_labels_np, 'pred_probabilities': current_best_val_preds_np})
  699. val_preds_df.to_csv(BEST_VAL_PREDS_SAVE_PATH, index=False) # Overwrites
  700. logging.info(f"Best validation set predictions (Epoch {epoch+1}) saved to {BEST_VAL_PREDS_SAVE_PATH}")
  701. metrics_ci_val_df = calculate_metrics_with_ci(current_best_val_labels_np, current_best_val_preds_np)
  702. metrics_ci_val_df.to_csv(METRICS_CI_VAL_SAVE_PATH, index=False) # Overwrites
  703. logging.info(f"Metrics with 95% CI for best validation set (Epoch {epoch+1}) saved to {METRICS_CI_VAL_SAVE_PATH}")
  704. logging.info(f"Best Validation Set Metrics (Epoch {epoch+1}, with 95% CI):\n" + metrics_ci_val_df.to_string())
  705. cm_data_val = np.array([[val_metrics_epoch.get('tn',0), val_metrics_epoch.get('fp',0)],
  706. [val_metrics_epoch.get('fn',0), val_metrics_epoch.get('tp',0)]])
  707. plt.figure(figsize=(8,6)); sns.heatmap(cm_data_val, annot=True, fmt='d', cmap='Blues', xticklabels=['Pred Neg', 'Pred Pos'], yticklabels=['Actual Neg', 'Actual Pos'])
  708. plt.title(f'Confusion Matrix (Best Val Epoch {epoch+1} - AUPRC: {best_val_auprc:.4f})'); plt.tight_layout()
  709. plt.savefig(FIG_PATH/f"confusion_matrix_best_val_epoch_current.png"); plt.close() # Overwrites
  710. logging.info(f"Confusion matrix for best validation epoch (Epoch {epoch+1}) saved to {FIG_PATH/'confusion_matrix_best_val_epoch_current.png'}")
  711. # --- Evaluate this new best model on Test Set and Save/Overwrite ---
  712. if test_dataloader:
  713. set_seed(seed) # Re-set seed for deterministic test pass
  714. model.eval() # Ensure model is in eval mode
  715. current_best_model_test_preds_list, current_best_model_test_labels_list = [], []
  716. test_eval_bar = tqdm(test_dataloader, desc=f"Epoch {epoch+1} [Best Model Test Eval]", leave=False)
  717. with torch.no_grad():
  718. for batch_test in test_eval_bar:
  719. ids_test, mask_test, tgts_test = batch_test['input_ids'].to(device), batch_test['attention_mask'].to(device), batch_test['labels'].to(device).unsqueeze(1)
  720. outputs_test = model(input_ids=ids_test, attention_mask=mask_test)
  721. current_best_model_test_preds_list.append(torch.sigmoid(outputs_test.detach()))
  722. current_best_model_test_labels_list.append(tgts_test.detach())
  723. current_best_model_test_preds_tensor = torch.cat(current_best_model_test_preds_list)
  724. current_best_model_test_labels_tensor = torch.cat(current_best_model_test_labels_list)
  725. current_best_model_test_preds_np = current_best_model_test_preds_tensor.cpu().numpy().flatten()
  726. current_best_model_test_labels_np = current_best_model_test_labels_tensor.cpu().numpy().flatten()
  727. test_preds_df_current_best = pd.DataFrame({'true_labels': current_best_model_test_labels_np, 'pred_probabilities': current_best_model_test_preds_np})
  728. test_preds_df_current_best.to_csv(BEST_TEST_PREDS_SAVE_PATH, index=False) # Overwrites
  729. logging.info(f"Test set predictions for current best model (from Val Epoch {epoch+1}) saved to {BEST_TEST_PREDS_SAVE_PATH}")
  730. metrics_ci_test_df_current_best = calculate_metrics_with_ci(current_best_model_test_labels_np, current_best_model_test_preds_np)
  731. metrics_ci_test_df_current_best.to_csv(METRICS_CI_TEST_SAVE_PATH, index=False) # Overwrites
  732. logging.info(f"Metrics with 95% CI for test set (current best model from Val Epoch {epoch+1}) saved to {METRICS_CI_TEST_SAVE_PATH}")
  733. logging.info(f"Current Best Model (from Val Epoch {epoch+1}) Test Set Metrics (with 95% CI):\n" + metrics_ci_test_df_current_best.to_string())
  734. current_best_model_test_metrics_basic = compute_metrics(current_best_model_test_preds_tensor, current_best_model_test_labels_tensor)
  735. cm_data_test_current_best = np.array([[current_best_model_test_metrics_basic.get('tn',0), current_best_model_test_metrics_basic.get('fp',0)],
  736. [current_best_model_test_metrics_basic.get('fn',0), current_best_model_test_metrics_basic.get('tp',0)]])
  737. plt.figure(figsize=(8,6)); sns.heatmap(cm_data_test_current_best, annot=True, fmt='d', cmap='Greens', xticklabels=['Pred Neg', 'Pred Pos'], yticklabels=['Actual Neg', 'Actual Pos'])
  738. auprc_test_display_current_best = current_best_model_test_metrics_basic.get('auprc', float('nan'))
  739. plt.title(f'CM Test (Best Val Ep {epoch+1} - AUPRC: {auprc_test_display_current_best:.4f})'); plt.tight_layout()
  740. plt.savefig(FIG_PATH/f"confusion_matrix_current_best_model_test.png"); plt.close() # Overwrites
  741. logging.info(f"Confusion matrix for current best model on test set saved to {FIG_PATH/'confusion_matrix_current_best_model_test.png'}")
  742. else: # AUPRC did not improve
  743. epochs_no_improve += 1
  744. if early_stopping_patience > 0 and epochs_no_improve >= early_stopping_patience:
  745. logging.info(f"Early stopping at Epoch {epoch+1} (Val AUPRC no improvement for {early_stopping_patience} epochs).")
  746. else: # No validation dataloader
  747. logging.info(f"Epoch {epoch+1} completed. No validation set for early stopping.")
  748. # End-of-Epoch Test Set Evaluation (always performed if test_dataloader exists)
  749. if test_dataloader:
  750. model.eval(); total_test_loss_epoch = 0; all_test_preds_epoch_eoe, all_test_labels_epoch_eoe = [], []
  751. test_progress_bar_eoe = tqdm(test_dataloader, desc=f"Epoch {epoch+1} [EOE Test]", leave=False) # EOE = End Of Epoch
  752. with torch.no_grad():
  753. for batch in test_progress_bar_eoe:
  754. ids, mask, tgts = batch['input_ids'].to(device), batch['attention_mask'].to(device), batch['labels'].to(device).unsqueeze(1)
  755. outputs = model(input_ids=ids, attention_mask=mask)
  756. loss = criterion(outputs, tgts) # Use same criterion for test loss
  757. total_test_loss_epoch += loss.item()
  758. all_test_preds_epoch_eoe.append(torch.sigmoid(outputs.detach()))
  759. all_test_labels_epoch_eoe.append(tgts.detach())
  760. test_progress_bar_eoe.set_postfix({'test_loss': f'{loss.item():.4f}'})
  761. avg_test_loss_epoch = total_test_loss_epoch / len(test_dataloader) if len(test_dataloader) > 0 else 0
  762. test_metrics_epoch_eoe = compute_metrics(torch.cat(all_test_preds_epoch_eoe), torch.cat(all_test_labels_epoch_eoe))
  763. current_epoch_metrics.update({'avg_test_loss': avg_test_loss_epoch, **{f'test_{k}': v for k,v in test_metrics_epoch_eoe.items()}})
  764. logging.info(f"Epoch {epoch+1} (End of Epoch Eval) - Test Loss: {avg_test_loss_epoch:.4f}, Test F1: {test_metrics_epoch_eoe.get('f1',0):.4f}, Test AUPRC: {test_metrics_epoch_eoe.get('auprc',0):.4f}")
  765. all_epoch_metrics_run.append(current_epoch_metrics)
  766. if val_dataloader and early_stopping_patience > 0 and epochs_no_improve >= early_stopping_patience:
  767. break # Break from training loop due to early stopping
  768. # --- After Training Loop ---
  769. # If no validation was done, save the model from the last trained epoch.
  770. # If validation was done, the best model (weights & LoRA) was already saved/overwritten during the loop.
  771. # This block handles the case where no validation is performed.
  772. if not val_dataloader and model: # Only save last epoch if no validation was performed at all
  773. logging.info("No validation set was used. Saving model from the last trained epoch.")
  774. last_epoch_state_to_save = {'fusion_module_state_dict': model.fusion_module.state_dict(), 'classifier_head_state_dict': model.classifier.state_dict()}
  775. torch.save(last_epoch_state_to_save, fusion_model_save_path)
  776. logging.info(f"Model from last trained epoch saved to {fusion_model_save_path}")
  777. if use_lora and PEFT_AVAILABLE:
  778. model.base_model.save_pretrained(str(lora_adapter_save_path))
  779. logging.info(f"LoRA adapters from last trained epoch saved to {lora_adapter_save_path}")
  780. elif not best_model_state_in_memory and model : # Should not happen if val_dataloader exists and runs for at least one epoch
  781. logging.warning("A best model was not identified during validation, but training completed. Saving model from last epoch as fallback.")
  782. last_epoch_state_to_save = {'fusion_module_state_dict': model.fusion_module.state_dict(), 'classifier_head_state_dict': model.classifier.state_dict()}
  783. torch.save(last_epoch_state_to_save, fusion_model_save_path)
  784. if use_lora and PEFT_AVAILABLE: model.base_model.save_pretrained(str(lora_adapter_save_path))
  785. elif not model:
  786. logging.warning("No model object available after training loop. Nothing to save.")
  787. if all_epoch_metrics_run:
  788. try:
  789. def convert_nan(o): return None if isinstance(o, float) and np.isnan(o) else o # JSON cannot serialize NaN
  790. with open(EPOCH_METRICS_SAVE_PATH, 'w') as f:
  791. json.dump(all_epoch_metrics_run, f, indent=4, default=convert_nan)
  792. logging.info(f"All epoch metrics saved to {EPOCH_METRICS_SAVE_PATH}")
  793. except Exception as e: logging.error(f"Failed to save epoch metrics: {e}")
  794. # ───── Text Embedding ───────────────────────────────────
  795. def embed_text_column(series: pd.Series, model_name: str = EMBED_MODEL, batch_size: int = 16,
  796. max_length: int = MAX_SEQ_LENGTH, num_layers_to_fuse: int = NUM_FUSION_LAYERS,
  797. trained_fusion_and_classifier_weights_path: Optional[Path] = None, # Path to .pth for fusion/classifier
  798. trained_lora_adapter_weights_path: Optional[Path] = None # Path to LoRA adapter directory
  799. ) -> pd.DataFrame:
  800. # Use global paths if specific paths are not provided
  801. actual_fusion_weights_path = TRAINED_FUSION_PATH if trained_fusion_and_classifier_weights_path is None else trained_fusion_and_classifier_weights_path
  802. actual_lora_adapter_path = TRAINED_LORA_ADAPTER_PATH if trained_lora_adapter_weights_path is None else trained_lora_adapter_weights_path
  803. device = "cuda" if torch.cuda.is_available() else "cpu"
  804. logging.info(f"Loading base model {model_name} to {device} for embedding generation...")
  805. base_model_for_embedding = AutoModel.from_pretrained(model_name, output_hidden_states=True)
  806. # Load LoRA adapters if available and path exists
  807. if actual_lora_adapter_path and actual_lora_adapter_path.exists() and PEFT_AVAILABLE:
  808. try:
  809. if actual_lora_adapter_path.is_dir() and any(actual_lora_adapter_path.iterdir()): # Check if dir and not empty
  810. base_model_for_embedding = PeftModel.from_pretrained(base_model_for_embedding, str(actual_lora_adapter_path))
  811. logging.info(f"Loaded LoRA adapters from {actual_lora_adapter_path} for embedding.")
  812. elif not actual_lora_adapter_path.is_dir():
  813. logging.warning(f"LoRA path {actual_lora_adapter_path} for embedding is not a directory. Using base model without LoRA.")
  814. else: # Is a directory but empty
  815. logging.warning(f"LoRA directory {actual_lora_adapter_path} for embedding is empty. Using base model without LoRA.")
  816. except Exception as e: logging.error(f"Error loading LoRA from {actual_lora_adapter_path} for embedding: {e}. Using base model.")
  817. elif actual_lora_adapter_path: logging.warning(f"LoRA path {actual_lora_adapter_path} for embedding not found, or PEFT unavailable. Using base model without LoRA.")
  818. base_model_for_embedding.to(device).eval()
  819. # Initialize and load fusion module weights
  820. layer_fusion_for_embedding = AdaptiveLayerFusion(base_model_for_embedding.config.hidden_size, num_layers_to_fuse, dropout=0.0).to(device) # Dropout 0 for inference
  821. if actual_fusion_weights_path and actual_fusion_weights_path.exists():
  822. try:
  823. checkpoint = torch.load(actual_fusion_weights_path, map_location=device)
  824. if 'fusion_module_state_dict' in checkpoint: # Check if it's the combined dict
  825. layer_fusion_for_embedding.load_state_dict(checkpoint['fusion_module_state_dict'])
  826. logging.info(f"Loaded fusion weights from 'fusion_module_state_dict' in {actual_fusion_weights_path} for embedding.")
  827. else: # Assume it's just the fusion module state dict
  828. layer_fusion_for_embedding.load_state_dict(checkpoint)
  829. logging.info(f"Loaded fusion weights directly from {actual_fusion_weights_path} for embedding (assumed no classifier dict).")
  830. except Exception as e: logging.error(f"Error loading fusion weights from {actual_fusion_weights_path} for embedding: {e}. Using initialized fusion module.")
  831. else: logging.warning(f"Fusion weights path {actual_fusion_weights_path} for embedding not found or None. Using initialized fusion module.")
  832. layer_fusion_for_embedding.eval()
  833. logging.info(f"Generating embeddings for {len(series)} texts, fusing top {num_layers_to_fuse} layers using loaded model...")
  834. tokenizer = AutoTokenizer.from_pretrained(model_name)
  835. texts_list = series.fillna("unknown").astype(str).tolist() # Handle NaNs and ensure string type
  836. all_embs, all_weights = [], []
  837. for i in tqdm(range(0, len(texts_list), batch_size), desc="Generating Embeddings", leave=False):
  838. batch_texts = texts_list[i:i+batch_size]
  839. try:
  840. encoded_input = tokenizer(batch_texts, padding=True, truncation=True, max_length=max_length, return_tensors='pt').to(device)
  841. with torch.no_grad():
  842. outputs = base_model_for_embedding(**encoded_input, output_hidden_states=True)
  843. hidden_states_all = list(outputs.hidden_states)
  844. # Validate num_layers_to_fuse for embedding model
  845. actual_num_tf_layers_emb = len(hidden_states_all) - 1 # Exclude initial embedding layer
  846. if not (0 < num_layers_to_fuse <= actual_num_tf_layers_emb):
  847. logging.warning(f"num_layers_to_fuse ({num_layers_to_fuse}) is invalid for embedding model with {actual_num_tf_layers_emb} TF layers. Clamping to [1, {actual_num_tf_layers_emb}]")
  848. fuse_n_emb = max(1, min(num_layers_to_fuse, actual_num_tf_layers_emb))
  849. if actual_num_tf_layers_emb == 0 : fuse_n_emb = 0 # Edge case: no transformer layers
  850. else:
  851. fuse_n_emb = num_layers_to_fuse
  852. if fuse_n_emb == 0: # If no layers to fuse (e.g. model has no transformer layers)
  853. logging.error("Embedding: No transformer layers to fuse from base model. Returning zeros.")
  854. # Create zero embeddings and uniform weights as fallback
  855. pooled_embs_batch = torch.zeros((len(batch_texts), base_model_for_embedding.config.hidden_size), device=device)
  856. weights_tensor_batch = torch.ones((len(batch_texts),1), device=device) # Dummy weights
  857. else:
  858. selected_hidden_states_emb = hidden_states_all[-fuse_n_emb:]
  859. pooled_embs_batch, weights_tensor_batch = layer_fusion_for_embedding(selected_hidden_states_emb, encoded_input['attention_mask'])
  860. all_embs.append(pooled_embs_batch.cpu().numpy())
  861. all_weights.append(weights_tensor_batch.cpu().numpy())
  862. except Exception as e:
  863. logging.error(f"Error in embedding generation batch {i//batch_size + 1}: {e}", exc_info=True)
  864. # Fallback for this batch: zero embeddings and uniform weights
  865. emb_dim_fallback = base_model_for_embedding.config.hidden_size
  866. fallback_num_layers = fuse_n_emb if 'fuse_n_emb' in locals() and fuse_n_emb > 0 else 1
  867. if fallback_num_layers == 0: fallback_num_layers = 1 # Ensure at least 1 for division
  868. all_embs.append(np.zeros((len(batch_texts), emb_dim_fallback)))
  869. all_weights.append(np.ones((len(batch_texts), fallback_num_layers)) / fallback_num_layers)
  870. if not all_embs: logging.warning("No embeddings were generated."); return pd.DataFrame()
  871. all_embs_arr = np.vstack(all_embs); all_weights_arr = np.vstack(all_weights)
  872. if all_embs_arr.shape[0] == 0: logging.warning("Embeddings array is empty after vstack."); return pd.DataFrame()
  873. # Normalize embeddings (L2 norm)
  874. norms = np.linalg.norm(all_embs_arr, axis=1, keepdims=True)
  875. norm_embs = all_embs_arr / np.maximum(norms, 1e-10) # Avoid division by zero
  876. emb_df = pd.DataFrame(norm_embs, columns=[f"text_emb_{j}" for j in range(norm_embs.shape[1])], index=series.index)
  877. res_df = emb_df
  878. if all_weights_arr.size > 0 and all_weights_arr.shape[1] > 0 : # Check if weights were actually produced and have columns
  879. weights_df = pd.DataFrame(all_weights_arr, columns=[f"layer_weight_{k}" for k in range(all_weights_arr.shape[1])], index=series.index)
  880. res_df = pd.concat([emb_df, weights_df], axis=1)
  881. logging.info(f"Average layer weights from embedding: {all_weights_arr.mean(axis=0)}")
  882. else: logging.warning("Layer weights array for embeddings was empty or had no columns.")
  883. logging.info(f"Text embedding generation complete. Final shape: {res_df.shape}")
  884. return res_df
  885. # ───── Visualization ────────────────────────────────────────────
  886. def quick_plots(df: pd.DataFrame, target_col_name: str):
  887. if target_col_name not in df.columns:
  888. logging.warning(f"Target column '{target_col_name}' not found in DataFrame for plotting."); return
  889. try:
  890. # Target distribution bar plot
  891. plt.figure(figsize=(6,4)); sns.countplot(x=df[target_col_name])
  892. plt.title(f"Distribution of {target_col_name}"); plt.tight_layout()
  893. plt.savefig(FIG_PATH/f"{target_col_name.replace('/','_')}_dist.png"); plt.close()
  894. # Target distribution pie chart
  895. plt.figure(figsize=(8,6)); counts = df[target_col_name].value_counts()
  896. if not counts.empty:
  897. plt.pie(counts, labels=counts.index, autopct='%1.1f%%', startangle=90, colors=['lightblue','salmon'])
  898. plt.title(f'{target_col_name} Pie Chart'); plt.axis('equal'); plt.tight_layout()
  899. plt.savefig(FIG_PATH/f"{target_col_name.replace('/','_')}_pie.png"); plt.close()
  900. else:
  901. logging.warning(f"No data for pie chart of {target_col_name}")
  902. logging.info(f"Target distribution plots saved to {FIG_PATH}")
  903. except Exception as e: logging.error(f"Error during target column plotting for '{target_col_name}': {e}", exc_info=True)
  904. # Missing values plot (top 20 features)
  905. miss_percent = df.isna().mean() * 100
  906. if not miss_percent.empty:
  907. miss_head = miss_percent[miss_percent > 0].sort_values(ascending=False).head(20)
  908. if not miss_head.empty:
  909. try:
  910. plt.figure(figsize=(10,8)); miss_head.plot(kind="barh")
  911. plt.title("Top 20 Features with Missing Values (%)"); plt.xlabel("Percentage Missing (%)"); plt.tight_layout()
  912. plt.savefig(FIG_PATH/"missing_values_top20.png"); plt.close()
  913. logging.info(f"Missing values plot saved to {FIG_PATH/'missing_values_top20.png'}")
  914. except Exception as e: logging.error(f"Error during missing values plot: {e}", exc_info=True)
  915. else: logging.info("No features with missing values to plot.")
  916. else: logging.info("DataFrame has no columns or no missing values to plot statistics for.")
  917. # ───── Main Process ────────────────────────────────────────────
  918. def main(args): # `args` comes from argparse
  919. # Make global variables modifiable within this function scope if they are reassigned
  920. global EMBED_MODEL, TEXT_COL, NUM_FUSION_LAYERS, INPUT_PATH, EMBED_PATH, TARGET_COL, NUM_EPOCHS, LEARNING_RATE, TRAIN_BATCH_SIZE, MIN_SAMPLES_FOR_TOP_LAYER_FINETUNING, EARLY_STOPPING_PATIENCE, VALIDATION_SPLIT_RATIO_DEFAULT, TRAINED_FUSION_PATH, TRAINED_LORA_ADAPTER_PATH, GLOBAL_SEED, USE_LORA, LORA_R, LORA_ALPHA, LORA_DROPOUT, LORA_TARGET_MODULES, FIG_PATH, EPOCH_METRICS_SAVE_PATH, LOG_FILE_PATH, HYPERPARAMS_FILE_PATH, RUN_SPECIFIC_DIR, BEST_VAL_PREDS_SAVE_PATH, METRICS_CI_VAL_SAVE_PATH, BEST_TEST_PREDS_SAVE_PATH, METRICS_CI_TEST_SAVE_PATH, N_BOOTSTRAPS_CI, TEST_SPLIT_RATIO_CONST
  921. set_seed(GLOBAL_SEED) # Set seed at the very beginning
  922. # --- Setup Run-Specific Output Directory and Paths ---
  923. timestamp_str = datetime.now().strftime("%Y%m%d_%H%M%S")
  924. sanitized_model_name = EMBED_MODEL.replace("/", "_").replace("-", "_") # Sanitize for directory name
  925. RUN_SPECIFIC_DIR = BASE_RESULTS_DIR / f"{sanitized_model_name}_{timestamp_str}"
  926. RUN_SPECIFIC_DIR.mkdir(parents=True, exist_ok=True)
  927. logging.info(f"All outputs will be saved to: {RUN_SPECIFIC_DIR.resolve()}")
  928. FIG_PATH = RUN_SPECIFIC_DIR / "figures"
  929. FIG_PATH.mkdir(parents=True, exist_ok=True)
  930. EMBED_PATH = RUN_SPECIFIC_DIR / args.output_csv_basename # Final output CSV
  931. TRAINED_FUSION_PATH = RUN_SPECIFIC_DIR / args.fusion_weights_basename # Fusion/classifier weights
  932. TRAINED_LORA_ADAPTER_PATH = RUN_SPECIFIC_DIR / args.lora_adapter_dir_basename # LoRA adapters directory
  933. if USE_LORA: TRAINED_LORA_ADAPTER_PATH.mkdir(parents=True, exist_ok=True) # Create LoRA dir if using LoRA
  934. EPOCH_METRICS_SAVE_PATH = FIG_PATH / "all_epoch_metrics.json" # Metrics per epoch
  935. LOG_FILE_PATH = RUN_SPECIFIC_DIR / "run_log.log" # File log
  936. HYPERPARAMS_FILE_PATH = RUN_SPECIFIC_DIR / "hyperparameters.json" # Hyperparameters log
  937. BEST_VAL_PREDS_SAVE_PATH = RUN_SPECIFIC_DIR / "best_validation_set_predictions.csv"
  938. METRICS_CI_VAL_SAVE_PATH = RUN_SPECIFIC_DIR / "best_validation_metrics_with_ci.csv"
  939. BEST_TEST_PREDS_SAVE_PATH = RUN_SPECIFIC_DIR / "best_model_test_set_predictions.csv"
  940. METRICS_CI_TEST_SAVE_PATH = RUN_SPECIFIC_DIR / "best_model_test_metrics_with_ci.csv"
  941. # --- Setup File Logging ---
  942. file_handler = logging.FileHandler(LOG_FILE_PATH)
  943. file_handler.setLevel(logging.INFO) # Or logging.DEBUG for more verbosity in file
  944. root_logger = logging.getLogger()
  945. if root_logger.hasHandlers() and root_logger.handlers[0].formatter: # Use existing formatter if available
  946. formatter = root_logger.handlers[0].formatter
  947. file_handler.setFormatter(formatter)
  948. else: # Fallback formatter
  949. fallback_formatter = logging.Formatter("%(asctime)s | %(levelname)s | %(message)s")
  950. file_handler.setFormatter(fallback_formatter)
  951. root_logger.addHandler(file_handler)
  952. logging.info(f"Logging to console and to file: {LOG_FILE_PATH.resolve()}")
  953. # --- Log Hyperparameters ---
  954. actual_validation_split_for_logging = args.validation_split_ratio if args.validation_split_ratio > 0 else 0.0
  955. hyperparameters_to_save = {
  956. "SCRIPT_VERSION": SCRIPT_VERSION, "TIMESTAMP": timestamp_str,
  957. "GLOBAL_SEED": GLOBAL_SEED, "EMBED_MODEL": EMBED_MODEL, "TEXT_COL": TEXT_COL, "TARGET_COL": TARGET_COL,
  958. "NUM_FUSION_LAYERS": NUM_FUSION_LAYERS, "NUM_EPOCHS": NUM_EPOCHS, "LEARNING_RATE": LEARNING_RATE,
  959. "TRAIN_BATCH_SIZE": TRAIN_BATCH_SIZE, "MAX_SEQ_LENGTH": MAX_SEQ_LENGTH,
  960. "MIN_SAMPLES_FOR_TOP_LAYER_FINETUNING": MIN_SAMPLES_FOR_TOP_LAYER_FINETUNING,
  961. "EARLY_STOPPING_PATIENCE": EARLY_STOPPING_PATIENCE,
  962. "TEST_SPLIT_RATIO_FIXED": TEST_SPLIT_RATIO_CONST,
  963. "VALIDATION_SPLIT_RATIO_FROM_REMAINDER_ARG": args.validation_split_ratio, # The arg value
  964. "EFFECTIVE_VALIDATION_SPLIT_OVERALL": (1.0 - TEST_SPLIT_RATIO_CONST) * actual_validation_split_for_logging if actual_validation_split_for_logging > 0 else 0.0,
  965. "EFFECTIVE_TRAIN_SPLIT_OVERALL": (1.0 - TEST_SPLIT_RATIO_CONST) * (1.0 - actual_validation_split_for_logging) if actual_validation_split_for_logging > 0 else (1.0 - TEST_SPLIT_RATIO_CONST),
  966. "N_BOOTSTRAPS_CI": N_BOOTSTRAPS_CI,
  967. "USE_LORA": USE_LORA, "LORA_R": LORA_R, "LORA_ALPHA": LORA_ALPHA, "LORA_DROPOUT": LORA_DROPOUT,
  968. "LORA_TARGET_MODULES": LORA_TARGET_MODULES,
  969. "INPUT_CSV_PATH": str(INPUT_PATH.resolve()), # Store resolved paths
  970. "RUN_SPECIFIC_OUTPUT_DIR": str(RUN_SPECIFIC_DIR.resolve()),
  971. "OUTPUT_CSV_BASENAME_ARG": args.output_csv_basename,
  972. "FUSION_WEIGHTS_BASENAME_ARG": args.fusion_weights_basename,
  973. "LORA_ADAPTER_DIR_BASENAME_ARG": args.lora_adapter_dir_basename,
  974. "EMBED_CSV_FULL_PATH": str(EMBED_PATH.resolve()),
  975. "FIGURES_DIR_FULL_PATH": str(FIG_PATH.resolve()),
  976. "TRAINED_FUSION_WEIGHTS_FULL_PATH": str(TRAINED_FUSION_PATH.resolve()),
  977. "TRAINED_LORA_ADAPTER_DIR_FULL_PATH": str(TRAINED_LORA_ADAPTER_PATH.resolve()),
  978. "EPOCH_METRICS_FULL_PATH": str(EPOCH_METRICS_SAVE_PATH.resolve()),
  979. "LOG_FILE_FULL_PATH": str(LOG_FILE_PATH.resolve()),
  980. "BEST_VAL_PREDS_SAVE_FULL_PATH": str(BEST_VAL_PREDS_SAVE_PATH.resolve()),
  981. "METRICS_CI_VAL_SAVE_FULL_PATH": str(METRICS_CI_VAL_SAVE_PATH.resolve()),
  982. "BEST_TEST_PREDS_SAVE_FULL_PATH": str(BEST_TEST_PREDS_SAVE_PATH.resolve()),
  983. "METRICS_CI_TEST_SAVE_FULL_PATH": str(METRICS_CI_TEST_SAVE_PATH.resolve()),
  984. }
  985. try:
  986. with open(HYPERPARAMS_FILE_PATH, 'w') as f:
  987. json.dump(hyperparameters_to_save, f, indent=4)
  988. logging.info(f"Hyperparameters saved to {HYPERPARAMS_FILE_PATH.resolve()}")
  989. except Exception as e:
  990. logging.error(f"Failed to save hyperparameters: {e}", exc_info=True)
  991. logging.info("================ Configuration (from globals after args) ================")
  992. for key, value in hyperparameters_to_save.items():
  993. if "PATH" not in key.upper() and "DIR" not in key.upper(): # Don't print full paths again here
  994. logging.info(f"{key}: {value}")
  995. logging.info(f"Run output directory (resolved): {RUN_SPECIFIC_DIR.resolve()}")
  996. logging.info("=======================================================================")
  997. if USE_LORA and not PEFT_AVAILABLE:
  998. logging.error("LoRA requested but PEFT library not available. Exiting.")
  999. return None # Or raise SystemExit
  1000. logging.info("Step 0: Initial data loading and preprocessing")
  1001. if not INPUT_PATH.exists():
  1002. logging.error(f"Input CSV {INPUT_PATH} not found."); return None
  1003. try:
  1004. df = pd.read_csv(INPUT_PATH)
  1005. except Exception as e:
  1006. logging.error(f"Error reading {INPUT_PATH}: {e}", exc_info=True); return None
  1007. logging.info(f"Initial df shape: {df.shape}")
  1008. if TARGET_COL not in df.columns or TEXT_COL not in df.columns:
  1009. logging.error(f"Missing target ('{TARGET_COL}') or text ('{TEXT_COL}') column."); return None
  1010. # Drop rows where essential target or text columns are NaN
  1011. df.dropna(subset=[TARGET_COL, TEXT_COL], inplace=True)
  1012. try:
  1013. df[TARGET_COL] = df[TARGET_COL].astype(int) # Ensure target is integer
  1014. except ValueError as e:
  1015. logging.error(f"Cannot convert target column '{TARGET_COL}' to int: {e}", exc_info=True); return None
  1016. if df.empty:
  1017. logging.error("DataFrame is empty after dropping NaNs from target/text columns."); return None
  1018. logging.info(f"Shape after NaN drop in target/text: {df.shape}. Target dist:\n{df[TARGET_COL].value_counts(normalize=True)}")
  1019. # Prepare data for the training pipeline
  1020. texts_for_training_pipeline, labels_for_training_pipeline = df[TEXT_COL].astype(str).tolist(), df[TARGET_COL].tolist()
  1021. if not texts_for_training_pipeline: # Should be redundant if df is not empty
  1022. logging.error("No text samples available for the training pipeline."); return None
  1023. # Prepare LoRA config for training function
  1024. lora_target_modules_list_main = [m.strip() for m in LORA_TARGET_MODULES.split(',') if m.strip()] if isinstance(LORA_TARGET_MODULES, str) else (LORA_TARGET_MODULES if isinstance(LORA_TARGET_MODULES, list) else [])
  1025. lora_config_dict_for_training = {"r": LORA_R, "lora_alpha": LORA_ALPHA, "lora_dropout": LORA_DROPOUT, "target_modules": lora_target_modules_list_main}
  1026. try:
  1027. # Test dataloader is now created within train_and_save_fusion_model
  1028. train_and_save_fusion_model(
  1029. texts=texts_for_training_pipeline,
  1030. labels=labels_for_training_pipeline,
  1031. base_model_name=EMBED_MODEL,
  1032. num_fusion_layers_to_train=NUM_FUSION_LAYERS,
  1033. fusion_model_save_path=TRAINED_FUSION_PATH, # For fusion/classifier weights
  1034. lora_adapter_save_path=TRAINED_LORA_ADAPTER_PATH, # For LoRA adapters
  1035. use_lora=USE_LORA,
  1036. lora_config_dict=lora_config_dict_for_training,
  1037. validation_split_ratio=args.validation_split_ratio, # From CLI args
  1038. tokenizer_max_len=MAX_SEQ_LENGTH,
  1039. train_batch_size=TRAIN_BATCH_SIZE,
  1040. num_epochs=NUM_EPOCHS,
  1041. learning_rate=LEARNING_RATE,
  1042. early_stopping_patience=EARLY_STOPPING_PATIENCE,
  1043. seed=GLOBAL_SEED,
  1044. test_dataloader_main=None # Explicitly None, will be created inside
  1045. )
  1046. except Exception as e:
  1047. logging.error(f"Critical error during training pipeline: {e}", exc_info=True)
  1048. logging.warning("Attempting to proceed to embedding generation with pre-existing/initialized weights due to training error.")
  1049. logging.info("Step 1: Data cleaning for final embedding DataFrame (using original full df)")
  1050. df_for_embedding_generation = df.copy() # Use the original df (after initial NaN drop) for embedding generation
  1051. # Drop specified columns, except target and always_keep columns
  1052. cols_to_drop_for_final_csv = [c for c in DROP_COLS if c != TARGET_COL and c in df_for_embedding_generation.columns]
  1053. if cols_to_drop_for_final_csv:
  1054. df_for_embedding_generation.drop(columns=cols_to_drop_for_final_csv, errors="ignore", inplace=True)
  1055. df_for_embedding_generation.dropna(axis=1, how='all', inplace=True) # Drop fully empty columns
  1056. logging.info(f"Shape of DataFrame prepared for embedding generation: {df_for_embedding_generation.shape}")
  1057. logging.info("Step 2: Text embedding generation for the output CSV")
  1058. if TEXT_COL in df_for_embedding_generation.columns:
  1059. lora_path_for_final_embed = None
  1060. if USE_LORA and TRAINED_LORA_ADAPTER_PATH.exists() and TRAINED_LORA_ADAPTER_PATH.is_dir() and any(TRAINED_LORA_ADAPTER_PATH.iterdir()):
  1061. lora_path_for_final_embed = TRAINED_LORA_ADAPTER_PATH
  1062. logging.info(f"Using saved LoRA adapters from {TRAINED_LORA_ADAPTER_PATH} for final embedding generation.")
  1063. elif USE_LORA:
  1064. logging.warning(f"LoRA was set to True, but adapters at {TRAINED_LORA_ADAPTER_PATH} are not found or empty for final embedding. Embedding will proceed without these specific LoRA weights.")
  1065. text_embeddings_df_final = embed_text_column(
  1066. df_for_embedding_generation[TEXT_COL].astype(str), # Ensure text column is string
  1067. model_name=EMBED_MODEL,
  1068. num_layers_to_fuse=NUM_FUSION_LAYERS,
  1069. trained_fusion_and_classifier_weights_path=TRAINED_FUSION_PATH, # Use the path to the saved .pth
  1070. trained_lora_adapter_weights_path=lora_path_for_final_embed # Use the path to the saved LoRA adapters
  1071. )
  1072. if not text_embeddings_df_final.empty:
  1073. if TEXT_COL in df_for_embedding_generation.columns: # Drop original text column
  1074. df_for_embedding_generation.drop(columns=[TEXT_COL], inplace=True)
  1075. df_for_embedding_generation = pd.concat([
  1076. df_for_embedding_generation.reset_index(drop=True), # Reset index for clean concat
  1077. text_embeddings_df_final.reset_index(drop=True)
  1078. ], axis=1)
  1079. logging.info(f"Shape of DataFrame after adding embeddings for final output: {df_for_embedding_generation.shape}")
  1080. else: logging.warning("Text embedding for final output CSV returned an empty DataFrame.")
  1081. else: logging.warning(f"Text column '{TEXT_COL}' not in DataFrame for final embedding generation. Skipping embedding step for output CSV.")
  1082. quick_plots(df_for_embedding_generation, TARGET_COL) # Generate plots on the final df
  1083. # Save the final DataFrame
  1084. EMBED_PATH.parent.mkdir(parents=True, exist_ok=True) # Ensure directory exists
  1085. df_for_embedding_generation.to_csv(EMBED_PATH, index=False)
  1086. logging.info(f"Final embedded data (potentially with features) saved to {EMBED_PATH.resolve()}")
  1087. logging.info(f"All results for this run are in: {RUN_SPECIFIC_DIR.resolve()}")
  1088. logging.info("Script finished successfully.")
  1089. return df_for_embedding_generation
  1090. if __name__ == "__main__":
  1091. import argparse
  1092. parser = argparse.ArgumentParser(description="MIMIC-IV Text Embedding with Adaptive Fusion & LoRA. Saves all outputs to a run-specific directory.")
  1093. # File/Path Arguments
  1094. parser.add_argument("--input-csv", type=str, default=str(INPUT_PATH), help=f"Path to input CSV. Default: {INPUT_PATH}")
  1095. parser.add_argument("--output-csv-basename", type=str, default="embedded_data.csv", help="Basename for the output CSV with embeddings. Default: embedded_data.csv")
  1096. parser.add_argument("--fusion-weights-basename", type=str, default="trained_model_weights.pth", help="Basename for trained fusion/classifier weights. Default: trained_model_weights.pth")
  1097. parser.add_argument("--lora-adapter-dir-basename", type=str, default="lora_adapters", help="Basename for LoRA adapters directory. Default: lora_adapters")
  1098. # Model & Data Column Arguments
  1099. parser.add_argument("--embed-model", type=str, default=EMBED_MODEL, help=f"Base HuggingFace model for embedding. Default: {EMBED_MODEL}")
  1100. parser.add_argument("--text-col", type=str, default=TEXT_COL, help=f"Name of the text column in the input CSV. Default: {TEXT_COL}")
  1101. parser.add_argument("--target-col", type=str, default=TARGET_COL, help=f"Name of the target column in the input CSV. Default: {TARGET_COL}")
  1102. # Training Hyperparameter Arguments
  1103. parser.add_argument("--num-layers", type=int, default=NUM_FUSION_LAYERS, help=f"Number of top transformer layers to fuse. Default: {NUM_FUSION_LAYERS}")
  1104. parser.add_argument("--epochs", type=int, default=NUM_EPOCHS, help=f"Number of training epochs. Default: {NUM_EPOCHS}")
  1105. parser.add_argument("--lr", type=float, default=LEARNING_RATE, help=f"Learning rate for training. Default: {LEARNING_RATE}")
  1106. parser.add_argument("--train-batch-size", type=int, default=TRAIN_BATCH_SIZE, help=f"Training batch size. Default: {TRAIN_BATCH_SIZE}")
  1107. parser.add_argument("--max-seq-length", type=int, default=MAX_SEQ_LENGTH, help=f"Maximum sequence length for tokenizer. Default: {MAX_SEQ_LENGTH}")
  1108. # Advanced Training Arguments
  1109. parser.add_argument("--min-samples-tune", type=int, default=MIN_SAMPLES_FOR_TOP_LAYER_FINETUNING, help=f"Minimum samples to fine-tune top layers of base model (if not using LoRA). Default: {MIN_SAMPLES_FOR_TOP_LAYER_FINETUNING}")
  1110. parser.add_argument("--early-stopping-patience", type=int, default=EARLY_STOPPING_PATIENCE, help=f"Patience for early stopping based on validation AUPRC. 0 to disable. Default: {EARLY_STOPPING_PATIENCE}")
  1111. parser.add_argument("--validation-split-ratio", type=float, default=VALIDATION_SPLIT_RATIO_DEFAULT,
  1112. help=f"Proportion of (Train+Val) data to use for validation (after {TEST_SPLIT_RATIO_CONST*100:.1f}% test split). 0 for no validation. Default: {VALIDATION_SPLIT_RATIO_DEFAULT:.4f} (for ~12.5% overall val).")
  1113. # Reproducibility & CI Arguments
  1114. parser.add_argument("--global-seed", type=int, default=GLOBAL_SEED, help=f"Global random seed for reproducibility. Default: {GLOBAL_SEED}")
  1115. parser.add_argument("--n-bootstraps-ci", type=int, default=N_BOOTSTRAPS_CI, help=f"Number of bootstrap samples for CI calculation. Default: {N_BOOTSTRAPS_CI}")
  1116. # LoRA Specific Arguments
  1117. parser.add_argument("--use-lora", action='store_true', default=USE_LORA_DEFAULT, help="Enable LoRA for fine-tuning the base model. Default: False")
  1118. parser.add_argument("--lora-r", type=int, default=LORA_R_DEFAULT, help=f"LoRA r (rank). Default: {LORA_R_DEFAULT}")
  1119. parser.add_argument("--lora-alpha", type=int, default=LORA_ALPHA_DEFAULT, help=f"LoRA alpha. Default: {LORA_ALPHA_DEFAULT}")
  1120. parser.add_argument("--lora-dropout", type=float, default=LORA_DROPOUT_DEFAULT, help=f"LoRA dropout. Default: {LORA_DROPOUT_DEFAULT}")
  1121. parser.add_argument("--lora-target-modules", type=str, default=LORA_TARGET_MODULES_DEFAULT, help=f"LoRA target modules (comma-separated, e.g., 'query,key,value'). Default: '{LORA_TARGET_MODULES_DEFAULT}'")
  1122. args = parser.parse_args()
  1123. # Update global variables from parsed arguments
  1124. INPUT_PATH = Path(args.input_csv)
  1125. EMBED_MODEL = args.embed_model
  1126. TEXT_COL = args.text_col
  1127. TARGET_COL = args.target_col
  1128. GLOBAL_SEED = args.global_seed
  1129. N_BOOTSTRAPS_CI = args.n_bootstraps_ci
  1130. NUM_FUSION_LAYERS = args.num_layers
  1131. NUM_EPOCHS = args.epochs
  1132. LEARNING_RATE = args.lr
  1133. TRAIN_BATCH_SIZE = args.train_batch_size
  1134. MAX_SEQ_LENGTH = args.max_seq_length
  1135. MIN_SAMPLES_FOR_TOP_LAYER_FINETUNING = args.min_samples_tune
  1136. EARLY_STOPPING_PATIENCE = args.early_stopping_patience
  1137. # VALIDATION_SPLIT_RATIO_DEFAULT is updated by args.validation_split_ratio in main() logic
  1138. USE_LORA = args.use_lora
  1139. LORA_R = args.lora_r
  1140. LORA_ALPHA = args.lora_alpha
  1141. LORA_DROPOUT = args.lora_dropout
  1142. LORA_TARGET_MODULES = args.lora_target_modules
  1143. # Validate validation_split_ratio
  1144. if not (0 <= args.validation_split_ratio < 1): # Must be in [0, 1)
  1145. logging.error(f"Validation split ratio (from remainder) must be [0, 1). Got: {args.validation_split_ratio}")
  1146. exit(1) # Or raise ValueError
  1147. if args.validation_split_ratio == 0:
  1148. logging.warning(f"Validation split ratio is 0. No validation set will be created from the train/val portion. Test set is still {TEST_SPLIT_RATIO_CONST*100:.1f}%. Early stopping disabled. Model from last epoch saved. Best model on test set will be based on this last epoch model.")
  1149. elif args.early_stopping_patience <= 0 and args.validation_split_ratio > 0 : # Val split exists but no patience
  1150. logging.warning("Validation split > 0 but early stopping patience <=0. Early stopping effectively disabled (will run for all epochs or until AUPRC improves, but best model choice still uses validation AUPRC).")
  1151. main(args)

train.py at commit f8d05ab, under Apache-2.0 · at the source

Overview

Authors: Han Wang1,2,3, Guoguang Lao1,4, Ruoyun He1, Ting Liu1, Hejiao Luo5, Changqi Qin6, Hongying Luo3, Yingqi Liu2,3, Junmin Huang3, Zihan Wei3, Lu Chen3, Yongzhi Xu3, Ziqian Bi7, Junhao Song8, Tianyang Wang9, Xin Liang Chia10, Xuanhe Hou11, Huafeng Liu3, Junfeng Hao3,12, Chunjie Tian1,3
  1. Department of Otorhinolaryngology, Affiliated Hospital of Guangdong Medical University, Zhanjiang 524000, China
  2. First Clinical College, Guangdong Medical University, Zhanjiang, China
  3. Guangdong Provincial Key Laboratory of Autophagy and Major Chronic Non-communicable Diseases, Institute of Nephrology, Affiliated Hospital of Guangdong Medical University, Zhanjiang 524001, China
  4. The First Dongguan Affiliated Hospital, Guangdong Medical University, Dongguan 523710, China
  5. Department of Critical Care Medicine, Affiliated Hospital of Guangdong Medical University, Zhanjiang, China
  6. CICU, The Seventh Affiliated Hospital, Sun Yat-sen University, Shenzhen, China
  7. Beijing University of Technology, Beijing 100124, China
  8. China Agricultural University, Beijing 100083, China
  9. Xi’an Jiaotong-Liverpool University, Suzhou 215123, China
  10. JTB Technology Corp., Tainan 741, Taiwan
  11. Department of Radiation Oncology (MAASTRO), GROW - Research Institute for Oncology and Reproduction, Maastricht University, Maastricht, the Netherlands
  12. Department of Family Medicine, Shengjing Hospital of China Medical University, Shenyang 110022, China
Journal: iScience, volume 29, issue 6, article 116135
Dates: received 15 October 2025; accepted 12 May 2026; published online 4 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1016/j.isci.2026.116135 · PMID 42305595 · PMCID PMC13266190 · OpenAlex W7163517376
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: human (organism), clinical / translational (subfield)
Methods: Connectivity, Spectral & time-frequency, Statistics, Machine learning
Keywords: Health sciences, Medicine, Medical specialty, Internal medicine, Intensive care medicine
Topic: ECG Monitoring and Analysis (Cardiology and Cardiovascular Medicine, Medicine), according to OpenAlex
Funding: National Natural Science Foundation of China (National Science Foundation of China) (82160215, 81660174); Affiliated Hospital of Guangdong Medical University (GCC20220013, GCC2022046); Joint Research Project of Liaoning Provincial Technology Plan; Applied Basic Research Program (2023JH2/101700312)
Citations: not cited yet (Europe PMC); 54 references in the paper

Abstract

Early identification of ICU patients at high mortality risk is essential for triage and timely intervention. We present adaptive layer fusion with intelligent attention (ALFIA), a modular architecture that jointly trains low-rank adaptation (LoRA) adapters and an adaptive layer-weighting mechanism to fuse multi-layer semantic features from a pretrained transformer backbone. ALFIA operates on text-encoded representations of the structured EHR data, in which tabular clinical variables (demographics, vital signs, laboratory values, and severity scores) are converted into standardized natural-language descriptions rather than processed as free-text clinical notes. Evaluated on the CriticalWindow-24 benchmark with MIMIC-IV and eICU cohorts, ALFIA achieves strong AUPRC while maintaining a balanced precision-recall profile. The learned embeddings can be further combined with gradient boosting (ALFIA-boost) or neural networks (ALFIA-nn) for additional gains. These findings demonstrate that text-encoded structured EHR data can support practical, generalizable early-warning models for ICU mortality risk stratification.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 13 matches between paragraphs and lines of code.

Hanziwww/ALFIA

License: Apache-2.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: f8d05ab707faa81769056530e09121c2cbe4ecb8, 6 June 2025
Languages: Python (2)
Size: 5 files, 2 scripts
Software Heritage: not archived
Found in: “Data and code availability”
Holds: README, license file
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (2 files), pandas (2 files), PyTorch (2 files), scikit-learn (2 files), Hugging Face Transformers (2 files), Matplotlib (1 file), seaborn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
4 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 2 scripts, each with its path and the digest of its content;
  • 13 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Data and code availability

• The CriticalWindow-24 (CW-24) benchmark dataset is publicly available at https://github.com/Hanziwww/CW-24 and archived on Zenodo: 15574378 (https://zenodo.org/records/15574378). • The MIMIC-IV dataset is available through PhysioNet at https://physionet.org/content/mimiciv/3.1/ upon the completion of required training and data use agreement, and the eICU Collaborative Research Database is available at https://physionet.org/content/eicu-crd/2.0/ following the same credentialing requirements. • All original code for model implementation, training, and analysis is available at https://github.com/Hanziwww/ALFIA. • Any additional information required to reanalyze the data reported in this paper is available from the lead contact upon request.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 20 authors, 5 keywords, 4 funders, 42 references.

Cite

This paper

Wang, H., Lao, G., He, R., Liu, T., Luo, H., Qin, C., Luo, H., Liu, Y., Huang, J., Wei, Z., Chen, L., Xu, Y., Bi, Z., Song, J., Wang, T., Chia, X. L., Hou, X., Liu, H., Hao, J., & Tian, C. (2026). Predicting ICU in-hospital mortality from text-encoded structured EHR data using adaptive transformer layer fusion. iScience, 29(6), 116135. https://doi.org/10.1016/j.isci.2026.116135

BibTeX

@article{wang2026predicting,
author = {Wang, Han and Lao, Guoguang and He, Ruoyun and Liu, Ting and Luo, Hejiao and Qin, Changqi and Luo, Hongying and Liu, Yingqi and Huang, Junmin and Wei, Zihan and Chen, Lu and Xu, Yongzhi and Bi, Ziqian and Song, Junhao and Wang, Tianyang and Chia, Xin Liang and Hou, Xuanhe and Liu, Huafeng and Hao, Junfeng and Tian, Chunjie},
title = {{Predicting ICU in-hospital mortality from text-encoded structured EHR data using adaptive transformer layer fusion}},
journal = {iScience},
year = {2026},
month = jun,
volume = {29},
number = {6},
pages = {116135},
publisher = {Elsevier},
issn = {2589-0042},
doi = {10.1016/j.isci.2026.116135},
url = {https://doi.org/10.1016/j.isci.2026.116135},
pmid = {42305595},
pmcid = {PMC13266190}
}

RIS

TY - JOUR
AU - Wang, Han
AU - Lao, Guoguang
AU - He, Ruoyun
AU - Liu, Ting
AU - Luo, Hejiao
AU - Qin, Changqi
AU - Luo, Hongying
AU - Liu, Yingqi
AU - Huang, Junmin
AU - Wei, Zihan
AU - Chen, Lu
AU - Xu, Yongzhi
AU - Bi, Ziqian
AU - Song, Junhao
AU - Wang, Tianyang
AU - Chia, Xin Liang
AU - Hou, Xuanhe
AU - Liu, Huafeng
AU - Hao, Junfeng
AU - Tian, Chunjie
TI - Predicting ICU in-hospital mortality from text-encoded structured EHR data using adaptive transformer layer fusion
T2 - iScience
J2 - iScience
PY - 2026
DA - 2026/06/04
VL - 29
IS - 6
SP - 116135
SN - 2589-0042
PB - Elsevier
DO - 10.1016/j.isci.2026.116135
UR - https://doi.org/10.1016/j.isci.2026.116135
LA - en
ER -

CSL-JSON

{
"id": "10.1016/j.isci.2026.116135",
"type": "article-journal",
"title": "Predicting ICU in-hospital mortality from text-encoded structured EHR data using adaptive transformer layer fusion",
"container-title": "iScience",
"author": [
{
"family": "Wang",
"given": "Han"
},
{
"family": "Lao",
"given": "Guoguang"
},
{
"family": "He",
"given": "Ruoyun"
},
{
"family": "Liu",
"given": "Ting"
},
{
"family": "Luo",
"given": "Hejiao"
},
{
"family": "Qin",
"given": "Changqi"
},
{
"family": "Luo",
"given": "Hongying"
},
{
"family": "Liu",
"given": "Yingqi"
},
{
"family": "Huang",
"given": "Junmin"
},
{
"family": "Wei",
"given": "Zihan"
},
{
"family": "Chen",
"given": "Lu"
},
{
"family": "Xu",
"given": "Yongzhi"
},
{
"family": "Bi",
"given": "Ziqian"
},
{
"family": "Song",
"given": "Junhao"
},
{
"family": "Wang",
"given": "Tianyang"
},
{
"family": "Chia",
"given": "Xin Liang"
},
{
"family": "Hou",
"given": "Xuanhe"
},
{
"family": "Liu",
"given": "Huafeng"
},
{
"family": "Hao",
"given": "Junfeng"
},
{
"family": "Tian",
"given": "Chunjie"
}
],
"container-title-short": "iScience",
"volume": "29",
"issue": "6",
"page": "116135",
"DOI": "10.1016/j.isci.2026.116135",
"PMID": "42305595",
"PMCID": "PMC13266190",
"ISSN": "2589-0042",
"publisher": "Elsevier",
"URL": "https://doi.org/10.1016/j.isci.2026.116135",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
4
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1038/s43856-026-01817-x [code]
Visual prompt engineering for multimodal and irregularly sampled medical data.
Journal: Communications medicine
In common: Hugging Face Transformers, PyTorch, seaborn, 4 other tools, physionet.org/content/mimiciv, clinical / translational, 3 references
[2] doi:10.3389/fneur.2026.1848730
Early prediction of incident delirium in traumatic brain injury: a multicenter validated and interpretable machine learning approach.
Journal: Frontiers in neurology
In common: physionet.org/content/eicu-crd, physionet.org/content/mimiciv, clinical / translational, 1 reference
[3] doi:10.3389/fncom.2026.1824898
Deep learning guided propofol ketamine dosing and inflammation trajectories in elderly burns.
Journal: Frontiers in computational neuroscience
In common: physionet.org/content/eicu-crd, physionet.org/content/mimiciv, 1 reference
[4] doi:10.1038/s41598-026-50791-w [code]
Hybrid Vi+ECNN framework for advanced ADHD diagnostic accuracy in medical imaging.
Journal: Scientific reports
In common: Hugging Face Transformers, PyTorch, seaborn, 4 other tools, clinical / translational, 1 reference
[5] doi:10.1093/braincomms/fcag253 [code]
Disease detection and classification in temporal lobe epilepsy: step-wise versus simultaneous AI decision models in a multisite neuroimaging study.
Journal: Brain communications
In common: Hugging Face Transformers, PyTorch, seaborn, 4 other tools, clinical / translational, 1 reference
[6] doi:10.1371/journal.pone.0354511 [code]
TokenUNet: A new case for transformers integration in efficient and interpretable 3D UNets for brain imaging segmentation.
Journal: PloS one
In common: Hugging Face Transformers, PyTorch, seaborn, 4 other tools, 1 reference
[7] doi:10.3389/fneur.2026.1753639
Early prophylactic heparin use is associated with reduced mortality in patients with non-traumatic subarachnoid hemorrhage.
Journal: Frontiers in neurology
In common: physionet.org/content/mimiciv, clinical / translational, 2 references
[8] doi:10.1038/s43856-026-01606-6 [code]
Validation of remote multimodal AI screening for Parkinson disease across diverse settings.
Journal: Communications medicine
In common: Hugging Face Transformers, PyTorch, seaborn, 4 other tools, clinical / translational
[9] doi:10.1038/s42003-026-10011-7 [code]
Learning brain dynamics across distinct scaling regimes reveals psychiatric signatures.
Journal: Communications biology
In common: Hugging Face Transformers, PyTorch, scikit-learn, 3 other tools, 1 reference
[10] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: Hugging Face Transformers, PyTorch, seaborn, 4 other tools

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.