Poda dinámica en inferencia
Rubén Rodríguez Abril
La poda dinámica en modelos de lenguaje es la desactivación, realizada en inferencia, de partes de un modelo en función de la tarea lingüística encomendada y de la naturaleza de los datos de entrada. Es esencial para evitar la sobrecarga del hardware. Engloba diferentes estrategias como poda de tokens, activaciones, cabezas de atención y capas, que permiten ahorrar costes computacionales y memoria sin pérdida significativa de precisión.
Introducción
Con el nombre de poda dinámica se designa a aquél conjunto de técnicas que apagan partes de un modelo durante la fase de inferencia, con el objetivo de reducir su gasto en poder computacional y memoria. Los pesos sinápticos y demás elementos fijos de la arquitectura no se modifican. En su lugar, son mecanismos específicos los que deciden, a partir de los datos de entrada, qué capas o subcapas desactivar.
Figura 1. Ejemplo de poda dinámica de capas en un modelo de lenguaje que está leyendo un texto. En cada nuevo token, las capas desactivadas (en rayas celestes) son diferentes. Fuente: Radial Networks: Dynamic Layer Routing for High-Performance Large Language Models.
La poda es una manifestación del principio de esparsión (sparsity), que implica que sólo un porcentaje reducido de los pesos sinápticos de una capa deben ser diferentes a cero. Este principio es el que subyace a las capas convolucionales de los sistemas de visión artificial y a los sistemas de mezcla de expertos. No sólo ahorra capacidad computacional y memoria, sino que también previene el sobreajuste (overfitting) de la red. Hay una cierta analogía entre el principio de esparsión y lo que sucede en el cerebro humano, en el que sólo un 1-4% de las neuronas están activas a la vez, lo que permite ahorrar energía, reducir el ruido y aumentar la especificidad de las señales.
En este artículo analizaremos cuatro tipos diferentes de poda: de tokens (o poda KV), activaciones, cabezas de atención y capas al completo. La primera consiste en la eliminación de parte de la información que fluye a través del modelo. La segunda, en el descarte de las activaciones que no sobrepasen un determinado umbral. Las dos últimas desactivan secciones enteras (capas o subcapas) de la arquitectura.
Poda dinámica de tokens (caché KV)
En el marco de los mecanismos de atención, la matriz de atención juega un papel fundamental. Sus pesos no son estáticos, sino dinámicos, ya que se forman a partir de la multiplicación matricial de Q (vectores consulta o query) y K (vectores clave). Tres transformaciones lineales convierten los datos de entrada en vectores consulta Q, vectores clave K y vectores valor V. Los dos primeros se usan para generar la matriz de atención en la forma descrita en la figura 2.
El tamaño de la matriz de atención es igual al cuadrado del número de vectores/tokens de entrada. Por consiguiente, a medida que la longitud de la secuencia de entrada crece linealmente, el tamaño de la matriz de atención aumenta cuadráticamente, con la consiguiente congestión computacional.
Figura 2. En el paso 8, la matriz de atención tiene 64 coeficientes. Cada coeficiente (n,m) se calcula multiplicando escalarmente la consulta Qn por la clave Km. Como el transformer es autorregresivo podemos deshacernos de los coeficientes anteriores (en verde), que ya fueron utilizados para calcular los tokens anteriores de salida. Para calcular el nuevo token sólo necesitamos procesar la fila de abajo de la matriz de atención (en rojo), utilizando todas las claves producidas hasta el momento pero tan sólo el nuevo vector consulta, Q8.
El caché KV
Sin embargo, el hecho de que los transformers sean modelos autorregresivos simplifica bastante las cosas: en el paso i, la unidad de atención sólo debe computar el nuevo token i. Los tokens anteriores, ya calculados, no se modifican. Como consecuencia, para calcular la nueva salida sólo es necesario multiplicar la matriz de valores V por la última fila de la matriz de atención. Esto algebraicamente se expresa con estas fórmulas:
Ai = QiKT
Oi = AiV
donde Oi es el token de salida en la posición i, Ai es la última fila de la matriz de atención (en rojo, en la Figura 2) y V la matriz de valores (de dimensión i x dk, donde dk es la dimensión del modelo, el tamaño de los tokens).
A medida que avanza la producción de la nueva cadena, es necesario calcular una nueva fila de la matriz Ai así como el nuevo token Oi. Para ello el modelo debe conservar todos los vectores K y V producidos hasta ese momento, pero tan sólo el vector Qi más reciente. Obsérvese que en las dos ecuaciones anteriores, K y V son matrices rectangulares (en azul), mientras que Qi es un simple vector (en rojo). Esta diferencia dimensional explica que las matrices clave y valor deban ser almacenadas en un denominado caché KV, de acceso casi inmediato para las unidades aritmético-lógicas encargadas de gestionar las operaciones de la unidad de atención.
Para aliviar la sobrecarga de este caché, se han desarrollado diferentes técnicas de poda (pruning) que suponen la eliminación de la memoria de ciertos vectores clave o valor (o, más técnicamente hablando, la anulación de sus componentes). Entre las más importantes, podemos citar las siguientes:
1) Poda aleatoria (Random pruning)
Se eliminan aleatoriamente claves y valores (que no son más que filas de matrices) del caché KV.
Este mecanismo es fácil de implementar, si bien se corre el riesgo de eliminar tokens con información importante.
2) Poda basada en la norma K (KNorm-based pruning)
Las claves con menor magnitud vectorial (norma L2) son eliminadas. Dado que un vector de norma reducida tiene poca importancia en los cálculos posteriores, su descarte apenas afecta a la precisión del modelo.
3) Poda de pesos de atención observados (Observed attention weights pruning)
Se eliminan aquellas claves que han recibido una menor atención media en la matriz. La atención media es calculada sumando los coeficientes de cada una de las claves, y dividiendo el resultado entre el número de coeficientes no nulo.
Sistema usado, con variantes, en InfLLM e HiP Attention.
4) Últimas 64 consultas (Last 64 queries attention pruning)
Se eliminan las claves y los valores que se generaran más de 64 pasos atrás.
Se basa en el principio de que en una conversación los oyentes humanos olvidan, en teoría, lo primero que han oído. Uno de los principales inconvenientes de este método es que a menudo el contexto amplio y la significación semántica de alto nivel de un texto son fijados al inicio del mismo (piénsese en novelas como el Quijote cuyos protagonistas y el tono de la obra son presentados en el primer capítulo).
Es utilizado en el modelo SnapKV.
5) Poda basada en atención esperada (Expected attention weights pruning)
En lugar de basarse en datos históricos, se realiza una predicción sobre la importancia que las claves y valores tendrán en el futuro en la configuración de la atención.
6) Retención del primer y último token (First-last token retention)
Se retienen los primeros y últimos tokens y se descarta el resto. Puede perder información crucial que aparezca en la parte intermedia del texto.
Usado en el modelo StreamingLLM, optimizado para la inferencia en tiempo real.
7) Poda jerárquica (Hierarchical pruning)
Se divide el texto en pedazos (chunks) y se escogen aquellos que tengan una mayor representatividad dentro de los mismos. Presenta el inconveniente de que la estrategia de clasificación puede ocasionar una sobrecarga del modelo.
Figura 3. Esquema del funcionamiento de InfiniteHip Attention, con poda jerárquica. Las claves se agrupan en varios pedazos (“chunks”) de similar tamaño. Aquellos pedazos cuyos tokens más representativos tengan una puntuación más baja son descartados. Fuente: InfiniteHiP Attention.
Poda de activaciones
En cada capa, las activaciones que no sobrepasen un cierto umbral son anuladas. Es el principio que subyace a la función de paso de Heaviside o a la función ReLU (que fijan el umbral en 0), entre otras.
Algunos modelos, como HDP (Hybrid Dynamic Pruning), aplican este principio a nivel de celdas: la matriz de atención es dividida en celdillas (“bloques”) de 2×2, cuyos coeficientes se suman. Aquellas que hayan recibido puntuaciones por debajo de un determinado umbral son eliminadas. En el caso de BERT-Base, la poda del 70% de estos bloques apenas supone una merma en la precisión de este modelo, con una reducción de tan sólo el 1%.
Algunos trabajos recientes proponen la creación de aceleradores (chips) especializados en la realización de tareas de poda. Varios ejemplos recientes son el coprocesador HDP, Energon o AccelTran. Este último divide las matrices a multiplicar en celdas (tilings), que son procesadas por separado. Un módulo, DynaTran, que determina cuáles de ellas se desactivan y no participan en la multiplicación. Según su diseñador,si este acelerador pasara de la fase de prototipo a producción, tendría 93K menos consumo energético que una Raspberry Pi y 330K más rendimiento en inferencia, lo que posibilitaría la ejecución de modelos de lenguaje en dispositivos IoT.
Figura 4. En el diagrama se reflejan las tres fases del proceso HDP: multiplicación matricial de Q y K, poda de la matriz resultante y aplicación de máscara. Fuente: Hybrid Dynamic Pruning: A Pathway to Efficient Transformer Inference.
Poda dinámica de cabezas de atención
Algunas de las cabezas de atención del modelo son desactivadas, particularmente aquellas cuya contribución al procesamiento de la información de entrada sea insignificante.
El proceso de poda se lleva a cabo tras el cálculo de los coeficientes de atención. Aquellas cabezas que hayan obtenido una puntuación global por debajo de un umbral pu son podadas. La poda se realiza antes de que tenga lugar la multiplicación por la matriz de valores V.
También usado en HPD (Hybrid Dynamic Pruning), entre otros.
Poda dinámica de capas
La poda dinámica de capas es la desactivación de determinadas capas del modelo en fase de inferencia.
Entre las estrategias más importantes podemos señalar aquellas que detienen la ejecución del modelo cuando la confianza en la respuesta es ya alta (baja entropía en la información) así como las que dinámicamente desactivan capas dependiendo de la nueva información de entrada recibida por el sistema.
Métodos basados en la disminución de niveles de entropía
La entropía de Shannon mide la incertidumbre de la información que fluye a través del modelo. La inferencia se detiene (early exit, “salida temprana”) cuando la entropía desciende por debajo de un umbral mínimo τ, lo que implica que la confianza en la predicción es ya alta.
Para ello, en diferentes capas de la arquitectura se sitúan “puntos de salida”, que son lugares en los que se mide la entropía de la información que está atravesando ese lugar. Si la entropía baja del umbral τ, el modelo detiene su funcionamiento y ofrece como salida la de su capa intermedia, en lugar de procesar las capas restantes.
En un transformer BERT dotado de 12 capas y entrenado para clasificar correos como spam o no spam, la implementación de este método podría realizarse mediante el modo siguiente:
-En las capas 4 y 8 se colocan puntos de salida en los cuales se extrae el token [CLS], que se hace pasar por un clasificador auxiliar. Se fija el umbral de entropía τ en 0,5.
-Un correo electrónico es introducido en el modelo. En la capa 4 el clasificador devuelve el vector [0,6, 0,4], indicando un 60% de probabilidades de spam. Dado que la entropía del vector es 0,673 (superior al umbral señalado) el modelo aun no tiene confianza suficiente en su predicción y por lo tanto la información sigue fluyendo hacia capas ulteriores.
-En el punto de salida de la capa 8 en cambio, el vector producido es [0,9, 0,1], que indica un 90% de posibilidades de spam. La entropía calculada (0,325) cae por debajo del umbral τ. El modelo detiene su funcionamiento y devuelve el resultado.
Figura 5. El sistema de salida temprana fue introducido originalmente en el año 2017 por la red convolucional BranchyNet, que se describe en esta imagen. En la actualidad su uso ha extendido a variantes de importantes modelos de lenguaje como BERT (p.e. FastBERT) o T5. Fuente: BranchNet, Fast Inference via Early Exiting from Deep Neural Networks.
Poda de capas y enrutamiento dinámico
Cada capa viene acompañada de un mecanismo de compuerta cuya respuesta binaria determina si la capa se ejecuta o no. Este mecanismo suele ser un perceptrón multicapas o una red recurrente, cuya decisión (la impresión de un 1 o un 0) no es diferenciable, lo cual dificulta la retropropagación de los gradientes.
El modelo Radial Networks introduce una interesante novedad: un router central distribuye la información a lo largo de todo la red y determina cuál será la próxima capa en ejecutarse y en procesar el nuevo token de una cadena.
Cada vez que un token ha sido procesado por una capa, la salida de ésta vuelve al router, que la reenvía a otra capa. El router es una mini-red cuya salida, dotada de una función softmax, determina a qué próxima capa enviará la información. Algebraicamente se expresa así:
zt = Router(et)
pt = softmax(zt)
lt+1 = argmax(pt)
Sin embargo, este esquema plantea el riesgo potencial de que la información pueda ser reenviada de nuevo a capas anteriores, tlo que en la práctica convertiría al modelo en una red neuronal recurrente. Para solventar este problema puede incorporarse una máscara a la función softmax, restringiendo el enrutamiento e impidiendo la formación de bucles de retroalimentación.
Figura 6: En Radial Networks un router central recibe un token de entrada y lo va reenviando sucesivamente a diversas capas del modelo. Fuente: Radial Networks: Dynamic Layer Routing for High-Performance Large Language Models.
Conclusión
Repetidos estudios han mostrado que no es necesario utilizar toda la capacidad teórica de cálculo de un modelo de lenguaje para ejecutar una tarea. Las técnicas de poda dinámica permiten asignar ciertas capas y cabezas de atención a la realización de una determinada tarea lingüística en función de la naturaleza de los datos de entrada. El uso combinado de las técnicas señaladas en este artículo ha permitido un ahorro significativo en términos de cálculo y memoria en modelos como BERT-Base, lo que ha impulsado su creciente adopción en el ámbito de la computación en el borde.
Lecturas Recomendadas
– InfiniteHiP: Extending Language Model Context Up to 3 Million Tokens on a Single GPU.
– BranchyNet, Fast Inference via Early Exiting from Deep Neural Networks.
– Radial Networks: Dynamic Layer Routing for High-Performance Large Language Models.
– Hybrid Dynamic Pruning: A Pathway to Efficient Transformer Inference.
– AccelTran: A Sparsity-Aware Accelerator for Dynamic Inference with Transformers.







