""" Centralized configuration for ACL-LKNet. All hyperparameters, paths, and experiment switches are defined here. Designed for Google Colab free tier (T4 GPU, 15GB VRAM). """ from dataclasses import dataclass, field from typing import List, Optional import os @dataclass class Config: """Master configuration for ACL-LKNet training pipeline.""" # ── Experiment ────────────────────────────────────────────────── experiment_name: str = "acl_lknet_v1" seed: int = 42 # ── Paths (Colab defaults) ────────────────────────────────────── # These are overridden in the Colab notebook after Drive mount data_dir: str = "/content/mrnet" drive_dir: str = "/content/drive/MyDrive/ACL_LKNet" checkpoint_dir: str = "" # set in __post_init__ log_dir: str = "" # set in __post_init__ # ── Data ──────────────────────────────────────────────────────── img_size: int = 224 # Resize slices to this (ImageNet standard) max_slices: int = 24 # Subsample to this many slices per view num_workers: int = 2 # Colab has limited CPU pin_memory: bool = True # ── Backbone ──────────────────────────────────────────────────── backbone: str = "convnext_tiny" # 'convnext_tiny', 'resnet18', 'resnet50', 'efficientnet_b0' pretrained: bool = True # Load ImageNet pretrained weights feature_dim: int = 0 # Auto-detected from backbone grad_checkpoint: bool = False # False to prevent PyTorch 2.4+ CheckpointError; slice_chunk_size=8 keeps VRAM <1.5GB slice_chunk_size: int = 8 # Process this many slices at once through backbone # Large-kernel modification (for ablation) use_large_kernels: bool = False # Replace DW conv kernels with larger ones large_kernel_sizes: List[int] = field(default_factory=lambda: [7, 13, 21, 31]) # ── Slice Attention & Aggregation ────────────────────────────── aggregation: str = "attention" # 'attention', 'max', 'mean' attn_hidden_dim: int = 256 # ── Cross-View Fusion ─────────────────────────────────────────── fusion_type: str = "attention" # 'attention' or 'concat' fusion_num_heads: int = 2 fusion_dropout: float = 0.1 # ── Classification Head ───────────────────────────────────────── classifier_hidden: int = 256 classifier_dropout: float = 0.3 # ── SSL: Masked Slice Modeling ────────────────────────────────── ssl_enabled: bool = True mask_ratio: float = 0.5 mask_strategy: str = "random" # 'random', 'contiguous', 'structured', 'mixed' msm_decoder_layers: int = 2 msm_decoder_dim: int = 256 msm_decoder_heads: int = 4 ssl_lr: float = 1e-4 ssl_weight_decay: float = 1e-4 ssl_patience: int = 15 # Early stopping patience for SSL ssl_save_every: int = 10 # ── Supervised Training ───────────────────────────────────────── num_epochs: int = 40 # Maximum epochs (monitor-based early stopping) epochs: Optional[int] = None # Alias for num_epochs (compatibility) batch_size: int = 1 # 1 exam at a time (T4 memory) accumulation_steps: int = 8 # Effective batch = batch_size × accumulation_steps lr: float = 3e-4 # LR for new layers backbone_lr: float = 1e-5 # LR for pretrained backbone (differential) weight_decay: float = 1e-4 label_smoothing: float = 0.05 mixup_alpha: float = 0.0 # Disabled: Mixup requires B>=2 but batch_size=1 is mandated for T4 VRAM gradient_clip: float = 1.0 warmup_epochs: int = 5 patience: int = 20 # Early stopping on val AUROC save_every: int = 5 # ── Class Imbalance ───────────────────────────────────────────── # MRNet ACL: 23.3% positive → weight ≈ 3.3 for positive class pos_weight: float = 3.3 # ── Regularization ────────────────────────────────────────────── stochastic_depth: float = 0.1 ema_decay: float = 0.999 # ── Mixed Precision ───────────────────────────────────────────── use_amp: bool = True # ── Augmentation (anatomically justified) ─────────────────────── use_horizontal_flip: bool = False # DISABLED — changes L/R anatomy rotation_degrees: float = 10.0 # Slight positioning variation translate_range: float = 0.05 # Minor FOV variation scale_range: tuple = (0.95, 1.05) # Scanner variation brightness: float = 0.15 # MRI intensity variation contrast: float = 0.15 # Scanner contrast variation gaussian_blur_p: float = 0.3 # Resolution variation random_erasing_p: float = 0.15 # Artifact simulation # ── Evaluation & Advanced Metrics ─────────────────────────────── bootstrap_n: int = 1000 confidence_level: float = 0.95 n_splits: int = 5 # Stratified 5-Fold Cross-Validation eval_threshold_mode: str = "youden" # 'default', 'youden', 'f1_optimal', 'high_sensitivity' target_sensitivity: float = 0.95 # High-sensitivity screening threshold explainability_target_slice_range: tuple = (10, 18) # Central cruciate ligament slices in 24-slice volume def __post_init__(self): if self.epochs is not None: self.num_epochs = self.epochs else: self.epochs = self.num_epochs self.checkpoint_dir = os.path.join(self.drive_dir, "checkpoints", self.experiment_name) self.log_dir = os.path.join(self.drive_dir, "logs", self.experiment_name) # Auto-detect feature dimension from backbone _feat_dims = { "convnext_tiny": 768, "resnet18": 512, "resnet50": 2048, "efficientnet_b0": 1280, } if self.feature_dim == 0: self.feature_dim = _feat_dims.get(self.backbone, 768) def to_dict(self): """Serialize config to dict for checkpoint saving.""" d = {} for k, v in self.__dict__.items(): try: # Ensure JSON-serializable import json json.dumps(v) d[k] = v except (TypeError, ValueError): d[k] = str(v) return d @classmethod def from_dict(cls, d): """Reconstruct config from dict.""" valid_fields = {f.name for f in cls.__dataclass_fields__.values()} filtered = {k: v for k, v in d.items() if k in valid_fields} return cls(**filtered) def get_training_config_table(self) -> List[tuple]: """Return structured hyperparameter specifications with scientific rationales.""" return [ ("Backbone Architecture", str(self.backbone), "Large effective receptive field for elongated ACL structure"), ("Input Resolution", f"{self.img_size} x {self.img_size}", "Standard ImageNet pretraining resolution"), ("Volume Slice Count", str(self.max_slices), "Uniform volumetric sequence normalization"), ("Slice Chunk Size", str(self.slice_chunk_size), "GPU VRAM containment on T4/P100 hardware"), ("Batch Size (Per Step)", str(self.batch_size), "Enforces volumetric integrity without OOM"), ("Gradient Accumulation", str(self.accumulation_steps), f"Yields effective batch size of {self.batch_size * self.accumulation_steps}"), ("Optimizer", "AdamW", "Decoupled weight decay for stable transformer/conv training"), ("Backbone LR", f"{self.backbone_lr:.1e}", "Differential LR to prevent catastrophic forgetting"), ("Head LR", f"{self.lr:.1e}", "Faster convergence for newly initialized modules"), ("Weight Decay", f"{self.weight_decay:.1e}", "L2 regularization penalty"), ("Loss Function", f"BCEWithLogitsLoss (pos_weight={self.pos_weight})", "Counteracts 23.3% ACL tear class imbalance"), ("Mixed Precision", "AMP FP16", "Accelerates training and caps memory allocation"), ("Model Averaging", f"EMA (decay={self.ema_decay})", "Smooths optimization landscape and boosts test generalization"), ("Label Smoothing", str(self.label_smoothing), "Mitigates overconfidence on small clinical cohorts"), ("Horizontal Flip", str(self.use_horizontal_flip), "Strictly DISABLED to preserve knee left/right anatomical asymmetry"), ("Rotation", f"±{self.rotation_degrees}°", "Simulates patient knee rotation within RF coil"), ("Cross-Validation", f"{self.n_splits}-Fold Stratified", "Unbiased cohort stability evaluation"), ("Bootstrap Iterations", f"N = {self.bootstrap_n} (95% CI)", "Empirical statistical accountability standard"), ] def export_config_markdown(self, save_path: Optional[str] = None) -> str: """Export training configuration table in Markdown format.""" rows = self.get_training_config_table() lines = [ "| Parameter | Value | Scientific & Clinical Rationale |", "| :--- | :--- | :--- |", ] for param, val, rationale in rows: lines.append(f"| **{param}** | `{val}` | {rationale} |") md_text = "\n".join(lines) if save_path: os.makedirs(os.path.dirname(os.path.abspath(save_path)), exist_ok=True) with open(save_path, "w") as f: f.write(md_text + "\n") return md_text def export_config_latex(self, save_path: Optional[str] = None) -> str: """Export training configuration table as a publication-ready LaTeX booktabs table.""" rows = self.get_training_config_table() lines = [ r"\begin{table}[htbp]", r"\centering", r"\caption{ACL-LKNet Training Hyperparameters and Experimental Configuration.}", r"\label{tab:training_configuration}", r"\begin{tabular}{lll}", r"\toprule", r"\textbf{Parameter} & \textbf{Configured Value} & \textbf{Scientific / Clinical Rationale} \\", r"\midrule", ] for param, val, rationale in rows: clean_param = param.replace("_", r"\_") clean_val = val.replace("_", r"\_").replace("%", r"\%") clean_rat = rationale.replace("_", r"\_").replace("%", r"\%") lines.append(f"{clean_param} & {clean_val} & {clean_rat} \\\\") lines.extend([ r"\bottomrule", r"\end{tabular}", r"\end{table}", ]) latex_text = "\n".join(lines) if save_path: os.makedirs(os.path.dirname(os.path.abspath(save_path)), exist_ok=True) with open(save_path, "w") as f: f.write(latex_text + "\n") return latex_text