ACL-LKNet / src /config.py
shareefch1413's picture
Upload folder using huggingface_hub
00801a0 verified
Raw History Blame Contribute Delete
12 kB
"""
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