El impacto de la normalización en la arquitectura Transformer
Deep Learning
Transformers
Matemáticas
Durante los últimos meses me he puesto a estudiar el libro de Deep Learning de Christopher Bishop y Hugh Bishop. Estudiando la arquitectura Transformer (que me parece muy interesante por diversos principios matemáticos que dan mucho de qué hablar, pero lo dejaré para otros posts) me di cuenta de que al usar la función softmax en la capa de atención, es supremamente importante que los inputs sean normalizados, lo cual despertó mi curiosidad de por qué, y es el origen de este blog.
Cuando estamos hablando de atención, estamos hablando de un concepto donde (en el mundo de las palabras) capturamos la naturaleza de la semántica de cada una de las palabras para con las demás, en el contexto de una frase.
12.1.5 Scaled Self-Attention
Hay que hacer un pequeño refinamiento que debemos realizar a la capa self-attention. Recordemos que los gradientes de la función softmax se vuelven exponencialmente más pequeños para inputs de altas magnitudes.
Demostración
Para una entrada:
z=[z1,…,zk]
Tenemos:
P=Softmax(z)=∑j=1kezjezi
Si tomamos:
Pi=Softmax(zi)=∑j=1kezjezi
Entonces, para P1:
P1=ez1+⋯+ezkez1
Ahora queremos saber qué tanto cambia P1 con respecto a z1:
∂z1∂P1
En la ecuación anterior podemos reemplazar el denominador por:
S=ez1+⋯+ezk
Entonces:
P1=Sez1
o:
P1=ez1S−1
Ahora aplicamos la regla del producto:
∂z1∂P1=∂z1∂ez1S−1+ez1∂z1∂S−1
Teniendo en cuenta:
∂x∂ex=ex
y:
∂x∂x−1=−x−2
Tenemos:
∂z1∂P1=ez1S−1−ez1S−2∂z1∂S
Pero:
S=ez1+⋯+ezk
por lo tanto:
∂z1∂S=ez1
Reemplazando:
∂z1∂P1=ez1S−1−ez1S−2ez1
Entonces:
∂z1∂P1=Sez1−S2e2z1
Pero:
P1=Sez1
y:
P12=S2e2z1
Por lo tanto:
∂z1∂P1=P1−P12
Finalmente:
∂z1∂P1=P1(1−P1)
Ahora revisemos qué pasa.
Si:
P1=1
entonces:
∂z1∂P1=1(1−1)≈0
Si:
P1=0.99
entonces:
∂z1∂P1=0.99(1−0.99)=0.0099
Es decir, cuando P1 se acerca a 1, el gradiente se acerca a 0.
Ahora llevemos esto a self-attention.
En self-attention calculamos los scores:
QKT
Si los vectores tienen una magnitud grande, estos scores también pueden tener una magnitud grande.
Al pasar estos valores por softmax, podemos terminar con probabilidades muy cercanas a 0 y 1.
Por ejemplo:
P=[0.001,0.998,0.001]
En este punto la función está saturada y sus gradientes son pequeños.
Por esto hacemos un pequeño refinamiento:
Softmax(dkQKT)
Dividimos los scores por dk.
Así reducimos su magnitud antes de aplicar softmax.
Con esto evitamos que softmax se sature tan fácilmente.