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
| #!/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() | |