#!/usr/bin/env python3
"""Comprueba tu GPU, compara su velocidad con la CPU y estima qué modelos te caben.

Uso: python3 comprobar_gpu.py   (necesita PyTorch: uv pip install torch)
"""

import time


def comprobar_gpu():
    try:
        import torch
    except ImportError:
        print("PyTorch no está instalado. Instálalo con: uv pip install torch")
        return

    print("=== Comprobación de GPU ===\n")
    print(f"Versión de PyTorch: {torch.__version__}")
    print(f"CUDA disponible: {torch.cuda.is_available()}")

    if not torch.cuda.is_available():
        if getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
            print("GPU de Apple (MPS) disponible: úsala con device = 'mps'.")
        else:
            print("\nNo se detecta GPU. No pasa nada: la mayoría de lecciones funcionan en CPU.")
            print("Para las lecciones pesadas, usa Google Colab (tiene GPU gratis).")
        return

    props = torch.cuda.get_device_properties(0)
    vram_gb = props.total_memory / 1e9
    print(f"Versión de CUDA: {torch.version.cuda}")
    print(f"GPU: {torch.cuda.get_device_name(0)}")
    print(f"Memoria: {vram_gb:.1f} GB")
    print(f"Capacidad de cómputo: {props.major}.{props.minor}")

    print("\n=== CPU frente a GPU ===\n")
    tam = 4000
    a = torch.randn(tam, tam)
    b = torch.randn(tam, tam)

    inicio = time.time()
    _ = a @ b
    t_cpu = time.time() - inicio
    print(f"Multiplicación de matrices {tam}x{tam} en CPU: {t_cpu:.3f} s")

    a_gpu, b_gpu = a.to("cuda"), b.to("cuda")
    torch.cuda.synchronize()
    inicio = time.time()
    _ = a_gpu @ b_gpu
    torch.cuda.synchronize()  # las operaciones en GPU son asíncronas
    t_gpu = time.time() - inicio
    print(f"Multiplicación de matrices {tam}x{tam} en GPU: {t_gpu:.3f} s")
    print(f"Aceleración: {t_cpu / t_gpu:.0f}x")

    print("\n=== ¿Qué modelos te caben? (solo los pesos, con un 20 % de margen) ===\n")
    for precision, bytes_param in [("fp16", 2), ("int8", 1), ("int4", 0.5)]:
        max_params = vram_gb * 0.8 * 1e9 / bytes_param
        print(f"  {precision}: hasta ~{max_params / 1e9:.0f}.000 millones de parámetros")


if __name__ == "__main__":
    comprobar_gpu()
