Spaces:
Runtime error
Runtime error
File size: 2,136 Bytes
cd95461 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 | """
utils/model_cache.py
====================
Cacheo de modelos con st.cache_resource.
El modelo se carga una sola vez por sesión de Streamlit.
Sin esto, cada interacción del usuario recargaría el modelo
completo — inutilizable en CPU.
"""
import streamlit as st
import torch
import sys
import os
# Agregar src al path — funciona tanto en Docker (/src) como en HF Spaces (./src)
for src_path in ['/src', os.path.join(os.path.dirname(__file__), '../../src')]:
if os.path.exists(src_path) and src_path not in sys.path:
sys.path.insert(0, src_path)
from activation_utils import cargar_modelo, obtener_capas_conv
@st.cache_resource(show_spinner="Cargando modelo... (solo ocurre una vez)")
def get_modelo(nombre_modelo: str):
"""
Carga y cachea un modelo preentrenado.
@st.cache_resource garantiza que el modelo se carga una sola vez
por sesión aunque el usuario cambie otros parámetros.
Args:
nombre_modelo: 'alexnet' o 'resnet18'
Returns:
Tupla (modelo, metadatos)
"""
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
modelo, meta = cargar_modelo(nombre_modelo)
modelo = modelo.to(device)
modelo.eval()
return modelo, meta
def get_capas_disponibles(nombre_modelo: str) -> list:
"""
Devuelve las capas convolucionales disponibles para un modelo.
Args:
nombre_modelo: 'alexnet' o 'resnet18'
Returns:
Lista de nombres de capas
"""
modelo, meta = get_modelo(nombre_modelo)
return list(obtener_capas_conv(modelo).keys())
def get_device_info() -> dict:
"""
Devuelve información sobre el device disponible.
"""
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
info = {
"device": str(device),
"tiene_gpu": torch.cuda.is_available(),
"gpu_nombre": None,
"gpu_vram_gb": None,
}
if torch.cuda.is_available():
info["gpu_nombre"] = torch.cuda.get_device_name(0)
info["gpu_vram_gb"] = round(
torch.cuda.get_device_properties(0).total_memory / 1e9, 1
)
return info |