sleep-apnea-shhs-v2 / inference_sleep_apnea.py
Gastpm's picture
Upload inference_sleep_apnea.py with huggingface_hub
f790fa8 verified
Raw
History Blame Contribute Delete
7.42 kB
#!/usr/bin/env python3
"""
================================================================================
INFERENCIA / EVALUACIÓN - DETECTOR APNEA DEL SUEÑO
================================================================================
Carga un modelo entrenado y predice sobre nuevos archivos .mat de SHHS.
Genera AHI (apneas por hora) por paciente y conteo de eventos.
Uso:
python inference_sleep_apnea.py \
--data_dir /path/to/test/mats \
--model_path /app/output/best_model.pt \
--output_dir /app/inference_results \
--batch_size 64
================================================================================
"""
import os
import glob
import json
import argparse
from collections import defaultdict
from typing import Dict, List, Optional
import numpy as np
import scipy.io as sio
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
# Importar definiciones del script de entrenamiento
from train_sleep_apnea import (
Config, load_mat_file, parse_patient_id, zscore_normalize,
create_windows, SleepDataset, collate_fn, SleepApneaNetV2,
compute_ahi
)
def load_model(checkpoint_path: str, config: Config):
"""Carga modelo entrenado desde checkpoint."""
device = config.device
model = SleepApneaNetV2(config).to(device)
checkpoint = torch.load(checkpoint_path, map_location=device)
model.load_state_dict(checkpoint["model_state_dict"])
model.eval()
print(f"Modelo cargado desde {checkpoint_path}")
print(f" Época: {checkpoint.get('epoch', 'N/A')}")
print(f" Val F1: {checkpoint.get('val_f1', 'N/A')}")
return model
def predict_on_files(model, file_paths: List[str], config: Config, threshold: float = 0.5):
"""
Predice sobre lista de archivos .mat.
Retorna dict con predicciones y AHI por paciente.
"""
results = {}
patient_predictions = defaultdict(lambda: {"windows": [], "preds": [], "labels": []})
for fpath in file_paths:
data = load_mat_file(fpath)
if data is None:
continue
pid = parse_patient_id(fpath)
windows = create_windows(data, config)
if not windows:
continue
# Normalizar por sujeto (ya hecho en create_windows)
ds = SleepDataset(windows, config, augment=False)
loader = DataLoader(ds, batch_size=config.batch_size, shuffle=False,
num_workers=0, collate_fn=collate_fn, pin_memory=False)
all_probs = []
all_preds = []
with torch.no_grad():
for batch in loader:
if len(batch) == 3:
x, cwt, y = batch
x, cwt = x.to(config.device), cwt.to(config.device)
logits = model(x, cwt)
else:
x, y = batch
x = x.to(config.device)
logits = model(x)
probs = F.softmax(logits, dim=-1)[:, 1].cpu().numpy()
preds = (probs > threshold).astype(int)
all_probs.extend(probs.tolist())
all_preds.extend(preds.tolist())
# Guardar resultados por archivo
results[fpath] = {
"patient_id": pid,
"n_windows": len(windows),
"predicted_events": int(sum(all_preds)),
"true_events": int(sum([w["label"] for w in windows])),
"predictions": all_preds,
"probabilities": all_probs,
}
# Agregar a acumulador por paciente
patient_predictions[pid]["windows"].extend(windows)
patient_predictions[pid]["preds"].extend(all_preds)
patient_predictions[pid]["labels"].extend([w["label"] for w in windows])
# Calcular AHI por paciente
patient_ahi = {}
stride_min = config.stride / config.fs / 60.0
for pid, data in patient_predictions.items():
preds = np.array(data["preds"])
labels = np.array(data["labels"])
wins = data["windows"]
pred_ahi = compute_ahi(wins, preds, stride_min=stride_min)
true_ahi = compute_ahi(wins, labels, stride_min=stride_min)
patient_ahi[pid] = {
"n_windows": len(wins),
"predicted_events": int(preds.sum()),
"true_events": int(labels.sum()),
"predicted_ahi": round(float(pred_ahi), 2),
"true_ahi": round(float(true_ahi), 2),
"error_ahi": round(float(pred_ahi - true_ahi), 2),
}
return results, patient_ahi
def main():
parser = argparse.ArgumentParser(description="Inferencia Sleep Apnea")
parser.add_argument("--data_dir", type=str, required=True)
parser.add_argument("--model_path", type=str, required=True)
parser.add_argument("--output_dir", type=str, default="/app/inference_results")
parser.add_argument("--threshold", type=float, default=0.5)
parser.add_argument("--batch_size", type=int, default=64)
parser.add_argument("--use_cwt", action="store_true")
parser.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu")
args = parser.parse_args()
os.makedirs(args.output_dir, exist_ok=True)
# Configuración por defecto (debe coincidir con entrenamiento)
cfg = Config()
cfg.data_dir = args.data_dir
cfg.batch_size = args.batch_size
cfg.device = args.device
cfg.use_cwt = args.use_cwt
# Cargar modelo
model = load_model(args.model_path, cfg)
# Encontrar archivos
files = sorted(glob.glob(os.path.join(args.data_dir, "*.mat")))
print(f"\nArchivos encontrados: {len(files)}")
if not files:
print("ERROR: No se encontraron archivos .mat")
return
# Inferencia
print(f"\nEjecutando inferencia (threshold={args.threshold})...")
results, patient_ahi = predict_on_files(model, files, cfg, threshold=args.threshold)
# Guardar resultados detallados
with open(os.path.join(args.output_dir, "inference_results.json"), "w") as f:
json.dump(results, f, indent=2)
with open(os.path.join(args.output_dir, "patient_ahi.json"), "w") as f:
json.dump(patient_ahi, f, indent=2)
# Reporte resumido
print("\n" + "=" * 70)
print("RESULTADOS POR PACIENTE")
print("=" * 70)
print(f"{'Paciente':<15} {'Ventanas':>10} {'Pred_Evts':>10} {'True_Evts':>10} {'Pred_AHI':>10} {'True_AHI':>10} {'Error':>10}")
print("-" * 80)
errors = []
for pid in sorted(patient_ahi.keys()):
d = patient_ahi[pid]
err = d["error_ahi"]
errors.append(err)
print(f"{pid:<15} {d['n_windows']:>10} {d['predicted_events']:>10} {d['true_events']:>10} "
f"{d['predicted_ahi']:>10.2f} {d['true_ahi']:>10.2f} {err:>10.2f}")
if errors:
errors = np.array(errors)
print("\n" + "=" * 70)
print("ESTADÍSTICAS GLOBALES AHI")
print("=" * 70)
print(f" MAE: {np.abs(errors).mean():.2f}")
print(f" RMSE: {np.sqrt((errors**2).mean()):.2f}")
print(f" Bias: {errors.mean():.2f}")
print(f"\nResultados guardados en: {args.output_dir}")
print(f" - inference_results.json (detalle por archivo)")
print(f" - patient_ahi.json (resumen por paciente)")
if __name__ == "__main__":
main()