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