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
Download src/utils.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 8.68 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/utils.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/src/utils.py
-
curl -L -o utils.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/utils.py
8.68 kB
| """ | |
| Utility functions for ACL-LKNet. | |
| Includes: reproducibility seeding, EMA model, checkpoint save/load, | |
| logging helpers, and Google Drive integration. | |
| """ | |
| import os | |
| import copy | |
| import random | |
| import logging | |
| from typing import Dict, Any, Optional | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| # ββ Reproducibility βββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def set_seed(seed: int = 42): | |
| """Set all random seeds for reproducibility.""" | |
| random.seed(seed) | |
| np.random.seed(seed) | |
| torch.manual_seed(seed) | |
| torch.cuda.manual_seed_all(seed) | |
| torch.backends.cudnn.deterministic = True | |
| torch.backends.cudnn.benchmark = False | |
| os.environ["PYTHONHASHSEED"] = str(seed) | |
| def get_rng_states() -> Dict[str, Any]: | |
| """Capture all RNG states for exact checkpoint reproducibility.""" | |
| states = { | |
| "python": random.getstate(), | |
| "numpy": np.random.get_state(), | |
| "torch": torch.get_rng_state(), | |
| } | |
| if torch.cuda.is_available(): | |
| states["cuda"] = torch.cuda.get_rng_state_all() | |
| return states | |
| def set_rng_states(states: Dict[str, Any]): | |
| """Restore RNG states from checkpoint.""" | |
| random.setstate(states["python"]) | |
| np.random.set_state(states["numpy"]) | |
| torch.set_rng_state(states["torch"]) | |
| if "cuda" in states and torch.cuda.is_available(): | |
| torch.cuda.set_rng_state_all(states["cuda"]) | |
| # ββ Exponential Moving Average ββββββββββββββββββββββββββββββββββββββ | |
| class EMAModel: | |
| """ | |
| Exponential Moving Average of model parameters. | |
| Maintains a shadow copy of model weights that is updated as: | |
| shadow = decay * shadow + (1 - decay) * current | |
| Use the EMA model for evaluation β it typically generalizes better. | |
| """ | |
| def __init__(self, model: nn.Module, decay: float = 0.999): | |
| self.decay = decay | |
| self.shadow = copy.deepcopy(model) | |
| self.shadow.eval() | |
| for p in self.shadow.parameters(): | |
| p.requires_grad_(False) | |
| def update(self, model: nn.Module): | |
| """Update shadow weights with current model weights and buffers.""" | |
| for s_param, m_param in zip(self.shadow.parameters(), model.parameters()): | |
| s_param.data.mul_(self.decay).add_(m_param.data, alpha=1.0 - self.decay) | |
| # Sync BatchNorm running statistics (buffers are not EMA-averaged, | |
| # they should directly mirror the training model's batch statistics) | |
| for s_buf, m_buf in zip(self.shadow.buffers(), model.buffers()): | |
| s_buf.data.copy_(m_buf.data) | |
| def state_dict(self): | |
| return self.shadow.state_dict() | |
| def load_state_dict(self, state_dict): | |
| self.shadow.load_state_dict(state_dict) | |
| def eval_model(self) -> nn.Module: | |
| """Return the shadow model for evaluation.""" | |
| return self.shadow | |
| # ββ Checkpoint Management βββββββββββββββββββββββββββββββββββββββββββ | |
| def save_checkpoint( | |
| path: str, | |
| epoch: int, | |
| phase: str, | |
| model: nn.Module, | |
| optimizer: torch.optim.Optimizer, | |
| scheduler: Any, | |
| scaler: Optional[torch.amp.GradScaler], | |
| ema: Optional[EMAModel], | |
| best_metric: float, | |
| best_epoch: int, | |
| train_history: list, | |
| val_history: list, | |
| patience_counter: int, | |
| config: Any, | |
| ): | |
| """ | |
| Save a full training checkpoint to Google Drive. | |
| Captures everything needed to resume training exactly: | |
| model, optimizer, scheduler, AMP scaler, EMA, RNG states, histories. | |
| """ | |
| os.makedirs(os.path.dirname(path), exist_ok=True) | |
| checkpoint = { | |
| "epoch": epoch, | |
| "phase": phase, | |
| "model_state_dict": model.state_dict(), | |
| "optimizer_state_dict": optimizer.state_dict(), | |
| "scheduler_state_dict": scheduler.state_dict() if scheduler else None, | |
| "scaler_state_dict": scaler.state_dict() if scaler else None, | |
| "ema_state_dict": ema.state_dict() if ema else None, | |
| "best_metric": best_metric, | |
| "best_epoch": best_epoch, | |
| "train_history": train_history, | |
| "val_history": val_history, | |
| "patience_counter": patience_counter, | |
| "rng_states": get_rng_states(), | |
| "config": config.to_dict() if hasattr(config, "to_dict") else str(config), | |
| } | |
| torch.save(checkpoint, path) | |
| logging.info(f"Checkpoint saved: {path}") | |
| def load_checkpoint( | |
| path: str, | |
| model: nn.Module, | |
| optimizer: Optional[torch.optim.Optimizer] = None, | |
| scheduler: Any = None, | |
| scaler: Optional[torch.amp.GradScaler] = None, | |
| ema: Optional[EMAModel] = None, | |
| ) -> Dict[str, Any]: | |
| """ | |
| Load a checkpoint and restore all training state. | |
| Returns the checkpoint dict for extracting histories, epoch, etc. | |
| """ | |
| checkpoint = torch.load(path, map_location="cpu", weights_only=False) | |
| model.load_state_dict(checkpoint["model_state_dict"]) | |
| if optimizer and "optimizer_state_dict" in checkpoint: | |
| optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) | |
| if scheduler and checkpoint.get("scheduler_state_dict"): | |
| scheduler.load_state_dict(checkpoint["scheduler_state_dict"]) | |
| if scaler and checkpoint.get("scaler_state_dict"): | |
| scaler.load_state_dict(checkpoint["scaler_state_dict"]) | |
| if ema and checkpoint.get("ema_state_dict"): | |
| ema.load_state_dict(checkpoint["ema_state_dict"]) | |
| if "rng_states" in checkpoint: | |
| set_rng_states(checkpoint["rng_states"]) | |
| logging.info(f"Checkpoint loaded: {path} (epoch {checkpoint['epoch']})") | |
| return checkpoint | |
| def find_latest_checkpoint(checkpoint_dir: str, phase: str = "finetune") -> Optional[str]: | |
| """Find the latest checkpoint file in the checkpoint directory.""" | |
| if not os.path.exists(checkpoint_dir): | |
| return None | |
| checkpoints = [ | |
| f for f in os.listdir(checkpoint_dir) | |
| if f.startswith(f"{phase}_") and f.endswith(".pt") | |
| ] | |
| if not checkpoints: | |
| return None | |
| # Sort by epoch number | |
| def extract_epoch(fname): | |
| try: | |
| parts = fname.replace(".pt", "").split("_epoch") | |
| return int(parts[-1]) | |
| except (ValueError, IndexError): | |
| return -1 | |
| checkpoints.sort(key=extract_epoch) | |
| return os.path.join(checkpoint_dir, checkpoints[-1]) | |
| # ββ Logging βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def setup_logging(log_dir: Optional[str] = None, level=logging.INFO): | |
| """Configure logging to console and optionally to file.""" | |
| handlers = [logging.StreamHandler()] | |
| if log_dir: | |
| os.makedirs(log_dir, exist_ok=True) | |
| handlers.append(logging.FileHandler(os.path.join(log_dir, "training.log"))) | |
| logging.basicConfig( | |
| level=level, | |
| format="%(asctime)s [%(levelname)s] %(message)s", | |
| datefmt="%Y-%m-%d %H:%M:%S", | |
| handlers=handlers, | |
| force=True, | |
| ) | |
| # ββ Metrics Formatting βββββββββββββββββββββββββββββββββββββββββββββ | |
| def format_metrics(metrics: Dict[str, float]) -> str: | |
| """Format a metrics dict into a readable string.""" | |
| parts = [] | |
| for k, v in metrics.items(): | |
| if isinstance(v, float): | |
| parts.append(f"{k}: {v:.4f}") | |
| else: | |
| parts.append(f"{k}: {v}") | |
| return " | ".join(parts) | |
| # ββ Memory Utils ββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def get_gpu_memory_info() -> Dict[str, float]: | |
| """Get GPU memory usage in MB.""" | |
| if not torch.cuda.is_available(): | |
| return {"allocated_mb": 0, "reserved_mb": 0, "total_mb": 0} | |
| try: | |
| props = torch.cuda.get_device_properties(0) | |
| total = getattr(props, "total_memory", getattr(props, "total_mem", 0)) / 1024**2 | |
| return { | |
| "allocated_mb": torch.cuda.memory_allocated() / 1024**2, | |
| "reserved_mb": torch.cuda.memory_reserved() / 1024**2, | |
| "total_mb": total, | |
| } | |
| except Exception: | |
| return {"allocated_mb": 0, "reserved_mb": 0, "total_mb": 0} | |
| def clear_gpu_memory(): | |
| """Force GPU memory cleanup.""" | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| torch.cuda.synchronize() | |