#!/usr/bin/env python3
"""Kit de depuración y profiling para proyectos de IA.

Uso:  python3 herramientas_depuracion.py     (ejecuta una demostración de cada herramienta)
Importa lo que necesites en tus proyectos:  from herramientas_depuracion import Cronometro, fuga_de_datos
Las funciones para tensores solo se usan si tienes PyTorch instalado.
"""

import cProfile
import io
import logging
import math
import pstats
import time
import tracemalloc


# ---------------------------------------------------------------- tiempo
class Cronometro:
    """Mide cuánto tarda un bloque de código:  with Cronometro("carga de datos"): ..."""

    def __init__(self, nombre=""):
        self.nombre = nombre
        self.segundos = 0.0

    def __enter__(self):
        self.inicio = time.perf_counter()
        return self

    def __exit__(self, *args):
        self.segundos = time.perf_counter() - self.inicio
        print(f"[{self.nombre}] {self.segundos:.4f} s")


def perfilar(funcion, *args, top=8, **kwargs):
    """Ejecuta una función con cProfile y muestra las llamadas que más tiempo acumulan."""
    perfil = cProfile.Profile()
    perfil.enable()
    resultado = funcion(*args, **kwargs)
    perfil.disable()
    texto = io.StringIO()
    pstats.Stats(perfil, stream=texto).sort_stats("cumulative").print_stats(top)
    print(texto.getvalue())
    return resultado


def memoria_por_linea(funcion, *args, top=5, **kwargs):
    """Muestra las líneas que más memoria reservan durante la ejecución de una función."""
    tracemalloc.start()
    resultado = funcion(*args, **kwargs)
    instantanea = tracemalloc.take_snapshot()
    tracemalloc.stop()
    for estadistica in instantanea.statistics("lineno")[:top]:
        print(estadistica)
    return resultado


# ---------------------------------------------------------------- logging
def configurar_logging(archivo="entrenamiento.log", nivel=logging.INFO):
    """Logs con fecha y nivel, en pantalla y en un archivo."""
    logging.basicConfig(
        level=nivel,
        format="%(asctime)s [%(levelname)s] %(message)s",
        handlers=[logging.FileHandler(archivo), logging.StreamHandler()],
        force=True,
    )
    return logging.getLogger("entrenamiento")


# ---------------------------------------------------------------- datos
def valores_invalidos(valores):
    """Posiciones con NaN o infinito en una lista de números (p. ej. las pérdidas)."""
    return [i for i, v in enumerate(valores) if math.isnan(v) or math.isinf(v)]


def fuga_de_datos(ids_train, ids_test):
    """Ids presentes en train y en test, y porcentaje del test contaminado."""
    train, test = set(ids_train), set(ids_test)
    solapados = train & test
    porcentaje = round(len(solapados) / len(test) * 100, 1) if test else 0.0
    if solapados:
        print(f"FUGA DE DATOS: {len(solapados)} ejemplos ({porcentaje} % del test) están también en train")
    return solapados, porcentaje


# ---------------------------------------------------------------- tensores (PyTorch)
def resumen_tensor(nombre, tensor):
    """Forma, tipo, dispositivo, rango y NaN de un tensor, en una línea."""
    print(f"{nombre}: forma={tuple(tensor.shape)}, tipo={tensor.dtype}, dispositivo={tensor.device}, "
          f"min={tensor.min().item():.4f}, max={tensor.max().item():.4f}, "
          f"media={tensor.float().mean().item():.4f}, tiene_nan={tensor.isnan().any().item()}")


def gradientes_problematicos(modelo):
    """Lista los parámetros con gradientes NaN o infinitos."""
    problemas = []
    for nombre, parametro in modelo.named_parameters():
        if parametro.grad is not None and (parametro.grad.isnan().any() or parametro.grad.isinf().any()):
            problemas.append(nombre)
    return problemas


def comprobar_dispositivos(modelo, *tensores):
    """Avisa si algún tensor no está en el mismo dispositivo que el modelo."""
    dispositivo = next(modelo.parameters()).device
    for i, t in enumerate(tensores):
        if t.device != dispositivo:
            print(f"AVISO: el tensor {i} está en {t.device} y el modelo en {dispositivo}")


# ---------------------------------------------------------------- demostración
def _trabajo_pesado(n=200_000):
    datos = [i * i for i in range(n)]
    return sum(sorted(datos, reverse=True)[:100])


def _crear_datos(n=2_000):
    return [list(range(100)) for _ in range(n)]


def main():
    print("=== 1. Cronómetro ===")
    with Cronometro("trabajo pesado"):
        _trabajo_pesado()

    print("\n=== 2. Profiling con cProfile ===")
    perfilar(_trabajo_pesado, top=5)

    print("=== 3. Memoria por línea con tracemalloc ===")
    memoria_por_linea(_crear_datos, top=3)

    print("\n=== 4. Valores inválidos en las pérdidas ===")
    perdidas = [0.9, 0.7, float("nan"), 0.5, float("inf")]
    print(f"Pérdidas: {perdidas} -> posiciones con problemas: {valores_invalidos(perdidas)}")

    print("\n=== 5. Fuga de datos ===")
    fuga_de_datos(range(0, 800), range(780, 1000))

    print("\n=== 6. Logging ===")
    log = configurar_logging()
    log.info("Empieza el entrenamiento: lr=%.4f, batch_size=%d", 1e-3, 32)
    log.warning("Pico de pérdida: %.4f en el paso %d", 7.31, 120)
    print("(también se ha guardado en entrenamiento.log)")

    try:
        import torch
    except ImportError:
        print("\n(PyTorch no está instalado: se omiten las herramientas para tensores)")
        return
    print("\n=== 7. Tensores (PyTorch) ===")
    x = torch.randn(4, 3)
    x[0, 0] = float("nan")
    resumen_tensor("x", x)


if __name__ == "__main__":
    main()
