Files
enlace/enlace/model/attention.py
T
msaldain 4048936067 Corregir lo que reportó ruff, que nunca se había ejecutado
ruff estaba configurado en pyproject.toml desde el primer commit y jamás se
había corrido. Tenía 15 hallazgos. Dos importan más allá del estilo:

- zip() sin strict= trunca en silencio al más corto. En las comparaciones de
  lotes eso significa que un test podía pasar sin haber comparado todo. Donde
  los largos deben coincidir ahora es strict=True; donde difieren a propósito
  (pares consecutivos) queda strict=False, que documenta la intención.
- Un import sin usar delataba algo peor: BraveBackend se había escrito sin una
  sola prueba. Se agregan ocho, contra una respuesta con la forma que devuelve
  la API, incluidas la limpieza de etiquetas, el caso de límite de tasa —que
  tiene que distinguirse de 'no respondió'— y que la credencial viaje en la
  cabecera y nunca en la URL. El respaldo tiene que funcionar justo cuando el
  primario ya falló; merecía la misma cobertura.

El resto es orden de imports, collections.abc y líneas largas.
2026-07-28 07:25:51 -03:00

48 lines
1.6 KiB
Python

"""Selección del backend de atención.
Es el único punto del modelo que conoce diferencias entre placas. FlashAttention-2
exige Ampere (sm_80) o superior: en la RTX 2060 (Turing, sm_75) hay que usar el
backend mem_efficient, que también es O(n) en memoria pero más lento. El perfil de
hardware decide; el modelo no sabe en qué placa corre.
"""
from __future__ import annotations
from collections.abc import Iterator
from contextlib import contextmanager
import torch
from torch.nn.attention import SDPBackend, sdpa_kernel
_BACKENDS: dict[str, list[SDPBackend]] = {
# "auto" deja elegir a PyTorch: prueba flash, después mem_efficient, después
# la implementación matemática. Es lo correcto salvo que se quiera forzar
# una ruta concreta para medir o para evitar un kernel con bugs.
"auto": [SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH],
"flash": [SDPBackend.FLASH_ATTENTION],
"mem_efficient": [SDPBackend.EFFICIENT_ATTENTION],
"math": [SDPBackend.MATH],
}
@contextmanager
def attention_backend(name: str) -> Iterator[None]:
"""Fija el backend de SDPA dentro del bloque.
En CPU no hay backends alternativos que elegir, así que es un no-op: forzar
uno ahí solo produce advertencias inútiles durante los tests.
"""
if not torch.cuda.is_available():
yield
return
try:
backends = _BACKENDS[name]
except KeyError:
raise ValueError(
f"backend de atención desconocido: {name!r}. Válidos: {sorted(_BACKENDS)}."
) from None
with sdpa_kernel(backends):
yield