ACL-LKNet / train_cv.py
shareefch1413's picture
Upload folder using huggingface_hub
00801a0 verified
Raw
History Blame Contribute Delete
5.08 kB
#!/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()