Image Classification
timm
English
medical-imaging
knee-mri
acl-tear-detection
deep-learning
convnext
self-attention
masked-slice-modeling
radiology
orthopedics
Eval Results (legacy)
Instructions to use shareefch1413/ACL-LKNet with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- timm
How to use shareefch1413/ACL-LKNet with timm:
import timm model = timm.create_model("hf-hub:shareefch1413/ACL-LKNet", pretrained=True) - Notebooks
- Google Colab
- Kaggle
File size: 11,967 Bytes
00801a0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 | """
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
|