Optimización de la KV cache en LLMs

En un post pasado vimos los recursos de memoria para los LLMs en producción o inferencia y que la KV cache, que almacena temporalmente los vectores Clave (Key) y Valor (Value) generados durante la lectura de la secuencia, es junto con el almacenamiento del modelo, uno de los mayores consumidores de RAM en inferencia.

Para opmitizar el uso de la KV cache existen diversas técnicas:

  • PagedAttention: Organiza la memoria de la caché en bloques pequeños no contiguos, similar a la paginación de memoria en un sistema operativo, lo que elimina la fragmentación y permite alojar más solicitudes simultáneas.
  • Cuantización de la KV Cache: Almacenar los tensores K/V en INT8 o INT4 en lugar de FP16/BF16, reduciendo el uso de memoria hasta 2-4x con una pérdida de calidad generalmente pequeña. Hugging Face transformers lo soporta de forma nativa mediante cache_implementation="quantized" (backends quanto o HQQ).
  • Agrupación de Atenciones: Reducen el número de cabezas de key/value (a 1 en MQA, o a un grupo reducido en GQA) mientras se mantienen múltiples cabezas de query. Esto disminuye directamente el tamaño de la cache proporcionalmente al número de cabezas KV eliminadas. Es la técnica usada en la mayoría de LLMs modernos.
  • Compresión y Expulsión de Caché: Elimina o comprime los bloques de claves y valores menos relevantes o poco utilizados a lo largo de una conversación extensa para liberar espacio en tiempo real.
  • Caché de Prompts (Prompt Caching): Reutiliza los prefijos de texto repetidos o fijos (como instrucciones del sistema o bases de conocimiento) entre distintas consultas para evitar recalcularlos. Muy usado en servidores de inferencia (vLLM, SGLang).
  • Multi-Head Latent Attention (MLA): Introducida por DeepSeek-V2: comprime K y V en un vector latente de baja dimensión antes de guardarlo en cache, y lo reconstruye en el momento de la atención, logrando reducciones drásticas de memoria sin usar GQA/MQA.

Vamos a ver un ejemplo que compara la cache dinámica estándar (FP16/FP32) frente a la cache cuantizada (INT4/INT8) (implementada de manera nativa en la librería Transformers) durante la generación de texto con el modelo Qwen2.5-1.5B-Instruct, midiendo memoria GPU pico utilizada y tiempo de generación.

Una vez importadas las librerías necesarias y cargado el modelo, generamos un texto de 4329 tokens.

# Repetimos un texto base varias veces para simular un contexto largo
# (por ejemplo, un documento o historial de conversacion extenso) y asi
# forzar una KV cache de mayor tamano.
texto_base = (
    "La inteligencia artificial esta transformando la forma en que las empresas "
    "analizan datos, automatizan procesos y toman decisiones estrategicas. "
    "Los modelos de lenguaje de gran escala permiten generar texto, resumir "
    "documentos, traducir idiomas y responder preguntas complejas con una "
    "fluidez cada vez mayor. "
)
prompt = texto_base * 60 + "En los proximos anos, se espera que"

inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
print("Tokens de entrada:", inputs['input_ids'].shape[1])

Definimos la función que genera el texto con las diferentes precisiones de almacenamiento de los tensores de la KV y medimos la memoria.

def benchmark_generacion(cache_implementation, max_new_tokens=700):
    """Genera texto con una implementacion de cache determinada y mide
    memoria GPU pico y tiempo de generacion."""
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.reset_peak_memory_stats()

    gen_kwargs = dict(
        max_new_tokens=max_new_tokens,
        do_sample=False,
        num_beams=1,
    )

    # 'dynamic' = cache estandar en FP16/FP32 (comportamiento por defecto)
    # 'quantized' = cache comprimida en INT4 via backend HQQ (mas estable en
    if cache_implementation is not None:
        gen_kwargs["cache_implementation"] = cache_implementation
        if cache_implementation == "quantized":
            gen_kwargs["cache_config"] = {"backend": "hqq", "nbits": 4}

    start = time.time()
    with torch.no_grad():
        output = model.generate(**inputs, **gen_kwargs)
    elapsed = time.time() - start

    peak_mem_mb = None
    incremento_mb = None
    if torch.cuda.is_available():
        peak_mem_mb = torch.cuda.max_memory_allocated() / (1024 ** 2)
        if memoria_base_mb is not None:
            incremento_mb = peak_mem_mb - memoria_base_mb

    texto = tokenizer.decode(output[0], skip_special_tokens=True)
    return elapsed, peak_mem_mb, incremento_mb, texto

Llamamos a la función e imprimimos los resultados.

resultados = {}

configuraciones = {
    "cache_estandar (dynamic)": "dynamic",
    "cache_cuantizada (quantized)": "quantized",
}

for nombre, impl in configuraciones.items():
    tiempo, memoria, incremento, texto = benchmark_generacion(impl)
    resultados[nombre] = {
        "tiempo_s": round(tiempo, 2),
        "memoria_pico_mb": round(memoria, 2) if memoria is not None else None,
        "incremento_mb": round(incremento, 2) if incremento is not None else None,
    }
    print(f"\n=== {nombre} ===")
    print(f"Tiempo de generacion: {tiempo:.2f} s")
    if memoria is not None:
        print(f"Memoria GPU pico total: {memoria:.2f} MB")
        print(f"Incremento sobre los pesos del modelo (activaciones + KV cache): {incremento:.2f} MB")
    else:
        print("Memoria GPU pico: N/D (ejecutando en CPU, sin medicion CUDA)")
    print("Fragmento generado:", texto[len(prompt):][:150], "...")


print("\nResumen comparativo\n" + "-" * 60)
for nombre, datos in resultados.items():
    print(
        f"{nombre:35s} | tiempo: {datos['tiempo_s']:>6} s | "
        f"pico total: {datos['memoria_pico_mb']} MB | "
        f"incremento (cache+activaciones): {datos['incremento_mb']} MB"
    )

Vemos en los resultados que el ahorro de memoria debido a la cuantización es de unos 80 megabytes, que es pequeño respecto al total de memoria usada porque:

  • El número de tokens generado es pequeño.
  • La memoria total almacena los pesos del modelo y las activaciones del prefill (estados intermedios y representaciones numéricas que se generan dentro de las capas).
  • El modelo Qwen2.5-1.5B-Instruct usa Grouped-Query Attention y tiene solo dos cabezas de Key y Value.
=== cache_estandar (dynamic) ===
Tiempo de generacion: 43.17 s
Memoria GPU pico total: 5238.16 MB
Incremento sobre los pesos del modelo (activaciones + KV cache): 2284.03 MB

=== cache_cuantizada (quantized) ===
Tiempo de generacion: 45.43 s
Memoria GPU pico total: 5156.99 MB
Incremento sobre los pesos del modelo (activaciones + KV cache): 2202.85 MB

Resumen comparativo
------------------------------------------------------------
cache_estandar (dynamic)            | tiempo:  43.17 s | pico total: 5238.16 MB | incremento (cache+activaciones): 2284.03 MB
cache_cuantizada (quantized)        | tiempo:  45.43 s | pico total: 5156.99 MB | incremento (cache+activaciones): 2202.85 MB

A continuación, hemos generado una tabla con el cálculo teórico del tamaño de las KV caches (FP16 e INT14) por número de tokens con los datos del modelo (número de capas, cabezas de atención, dimensión por cabeza). Para el orden de un poco menos de 5000 tokens el ahorro es de 84 MB, muy próximo a lo que hemos medido en la implementación real.

Deja una respuesta

Orgullosamente ofrecido por WordPress | Tema: Baskerville 2 por Anders Noren.

Subir ↑