#!/usr/bin/env python3 """ ACL-LKNet Training CLI ====================== Standalone command-line entrypoint for training ACL-LKNet models. Supports single-fold training, full 5-fold stratified cross-validation, and Phase 1 self-supervised pretraining (Masked Slice Modeling). Usage Examples: # Train Fold 1 with default configuration: python train_cv.py --data_dir /path/to/mrnet --fold 1 # Train all 5 folds sequentially for cross-validation: python train_cv.py --data_dir /path/to/mrnet --cv # Run Phase 1 Masked Slice Modeling (SSL) pretraining: python train_cv.py --data_dir /path/to/mrnet --ssl_pretrain """ import os import sys import argparse import logging import torch # Ensure local package imports work seamlessly sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from src.config import Config from src.train import train_supervised, train_ssl, train_5fold_cross_validation from src.utils import set_seed, setup_logging def parse_args(): parser = argparse.ArgumentParser( description="Train ACL-LKNet (Phase 1 SSL Pretraining or Phase 2 Supervised Fine-Tuning)." ) parser.add_argument( "--data_dir", type=str, default="./data/mrnet", help="Path to Stanford MRNet dataset root directory (containing train, valid, and test folders)." ) parser.add_argument( "--output_dir", type=str, default="./checkpoints", help="Directory where model checkpoints and logs will be saved." ) parser.add_argument( "--backbone", type=str, default="convnext_tiny", choices=["convnext_tiny", "resnet18", "replknet"], help="Backbone architecture (ConvNeXt-Tiny recommended for large receptive field)." ) parser.add_argument( "--fold", type=int, default=1, choices=[1, 2, 3, 4, 5], help="Fold index to train (1-5) when not running full cross-validation." ) parser.add_argument( "--cv", action="store_true", help="Run complete 5-fold stratified cross-validation." ) parser.add_argument( "--ssl_pretrain", action="store_true", help="Run Phase 1 Masked Slice Modeling (MSM) self-supervised pretraining." ) parser.add_argument( "--ssl_checkpoint", type=str, default=None, help="Optional path to SSL-pretrained backbone checkpoint to initialize Phase 2 training." ) parser.add_argument( "--epochs", type=int, default=30, help="Maximum training epochs (monitored by early stopping patience=20)." ) parser.add_argument( "--lr_backbone", type=float, default=1.0e-5, help="Differential learning rate for pretrained backbone." ) parser.add_argument( "--lr_head", type=float, default=3.0e-4, help="Learning rate for newly initialized attention and classification heads." ) parser.add_argument( "--batch_size", type=int, default=1, help="Physical batch size (keep = 1 on 15-16 GB GPUs to prevent OOM)." ) parser.add_argument( "--accum_steps", type=int, default=8, help="Gradient accumulation steps to achieve effective batch size = 8." ) parser.add_argument( "--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu", help="Compute device ('cuda' or 'cpu')." ) parser.add_argument( "--seed", type=int, default=42, help="Random seed for reproducibility across PyTorch, NumPy, and Python." ) return parser.parse_args() def main(): args = parse_args() # Initialize master configuration config = Config( data_dir=args.data_dir, checkpoint_dir=args.output_dir, log_dir=os.path.join(args.output_dir, "logs"), backbone_name=args.backbone, epochs=args.epochs, backbone_lr=args.lr_backbone, lr=args.lr_head, batch_size=args.batch_size, accumulation_steps=args.accum_steps, device=args.device, seed=args.seed, ) os.makedirs(config.checkpoint_dir, exist_ok=True) os.makedirs(config.log_dir, exist_ok=True) setup_logging(config.log_dir) set_seed(config.seed) logging.info(f"Initialized ACL-LKNet with Backbone: {config.backbone_name} on {config.device}") if args.ssl_pretrain: logging.info("Starting Phase 1: Masked Slice Modeling (MSM) SSL Pretraining...") train_ssl(config) elif args.cv: logging.info(f"Starting 5-Fold Stratified Cross-Validation on MRNet...") results = train_5fold_cross_validation(config, ssl_checkpoint=args.ssl_checkpoint) print("\n" + "=" * 60) print("Cross-Validation Complete! Summary:") print(f"Mean Val AUROC: {results.get('mean_val_auc', 'N/A')}") print("=" * 60) else: logging.info(f"Starting Single Fold Training: Fold {args.fold}...") config.experiment_name = f"acl_lknet_fold{args.fold}" train_supervised(config, fold_idx=args.fold - 1, ssl_checkpoint=args.ssl_checkpoint) if __name__ == "__main__": main()