skinscreen / app /ml /model_loader.py
soleil-kamitto
feat: aceptar solo imágenes del catálogo ISIC en /predict
87e8df5
Raw
History Blame Contribute Delete
2.34 kB
import logging
from functools import lru_cache
import torch
from torch import nn
from torchvision.models import ResNet50_Weights, resnet50
from app.config import settings
from app.models.schemas import CLASS_LABELS
logger = logging.getLogger(__name__)
NUM_CLASSES = len(CLASS_LABELS)
def build_model() -> nn.Module:
"""ResNet50 preentrenada en ImageNet + cabeza de clasificacion del spec:
GAP -> Dense(256, ReLU) -> Dropout(0.3) -> Dense(7, Softmax via CrossEntropy)."""
try:
model = resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)
except Exception:
logger.exception(
"No se pudieron descargar los pesos ImageNet (sin red?). "
"Usando ResNet50 con pesos aleatorios."
)
model = resnet50(weights=None)
in_features = model.fc.in_features
model.fc = nn.Sequential(
nn.Linear(in_features, 256),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(256, NUM_CLASSES),
)
return model
@lru_cache(maxsize=1)
def get_model() -> tuple[nn.Module, torch.device, bool]:
"""Singleton del modelo. Devuelve (modelo, device, fine_tuned).
Si existe un checkpoint entrenado en settings.model_path_resolved, lo carga.
Si no, usa la cabeza con pesos aleatorios (modo scaffold/demo) — las
predicciones son estructuralmente validas pero no clinicamente significativas.
"""
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = build_model()
fine_tuned = False
checkpoint_path = settings.model_path_resolved
if checkpoint_path.exists():
try:
state_dict = torch.load(checkpoint_path, map_location=device)
model.load_state_dict(state_dict)
fine_tuned = True
logger.info("Modelo fine-tuneado cargado desde %s", checkpoint_path)
except Exception:
logger.exception(
"No se pudo cargar checkpoint en %s, usando cabeza sin entrenar", checkpoint_path
)
else:
logger.warning(
"No hay checkpoint entrenado en %s — usando ResNet50 ImageNet con cabeza "
"sin fine-tuning (modo scaffold). Correr ml/scripts/train.py para entrenar.",
checkpoint_path,
)
model.to(device)
model.eval()
return model, device, fine_tuned