Ves al contingut
9/30Capítol 9 de 30

Attention i el bloc transformer, derivats d’una mitjana

Parteix del resum més barat d’un context —la mitjana—, mesura com falla i deixa que la fórmula d’attention surti de la reparació.

En aquesta pàgina

Arribes aquí amb un tokenizer del capítol 7, una taula d’embedding del capítol 8 i l’objectiu que els acompanya: donats els tokens fins ara, assignar una probabilitat al següent.

El que falta és el mig. Per predir el token tt, el model necessita un vector que resumeixi tot el que hi ha abans, i res del que has construït en produeix un. L’embedding del token t1t-1 no ho és: això és un model de bigrames, i no pot saber que la frase començava amb una pregunta. Una concatenació de tots els embeddings anteriors tampoc no ho és: el seu nombre canvia a cada pas, i una matriu de pesos fixa no pot acceptar una entrada de longitud variable.

Així doncs: un vector de mida fixa que resumeix un nombre variable de vectors. Aquest és tot el problema, i attention és el que obtens en resoldre’l de la manera més mandrosa possible i després reparar les dues coses que es trenquen.

La resposta que tenia el camp, i per què no la construïm

Enllaç a la secció: La resposta que tenia el camp, i per què no la construïm

Del 1997 fins aproximadament el 2017, el resum era un estat recurrent: mantenir un vector h\mathbf{h} i actualitzar-lo a cada token, ht=f(ht1,xt)\mathbf{h}_t = f(\mathbf{h}_{t-1}, \mathbf{x}_t). Mida fixa, entrada variable, exactament la forma adequada.

Va fallar de tres maneres, i l’arquitectura d’aquest capítol respon a totes tres. Fer backpropagation a través de TT passos multiplica TT jacobians, de manera que el gradient s’esvaeix o explota: la malaltia que el capítol 5 va mesurar dins d’un únic node tanh\tanh. L’LSTM1 es va dissenyar exactament contra això i va empènyer el rang útil de desenes de passos a centenars, sense canviar el fet que la informació del token 5 arriba al token 500 només si sobreviu a 495 actualitzacions seqüencials. Tota la font havia de cabre en un sol vector: en la traducció seqüència-a-seqüència2, un encoder comprimeix l’entrada en el seu estat final. Bahdanau, Cho i Bengio van anomenar aquest coll d’ampolla i el van arreglar el 2014, tres anys abans del transformer, permetent que el decoder prengués una suma ponderada de tots els estats de l’encoder amb pesos que calculava ell mateix.3 Tot el que ve a continuació és aquesta idea, aplicada per una seqüència a si mateixa, amb la recurrència eliminada. I l’actualització és seqüencial per construcció: ht\mathbf{h}_t necessita ht1\mathbf{h}_{t-1}, i una GPU amb deu mil nuclis no hi pot fer res. L’arquitectura que va guanyar no és òbviament més intel·ligent; és aquella en què el pas car és una multiplicació de matrius.

L’altre biaix inductiu clàssic, la convolució —fer lliscar un filtre petit per tota l’entrada, de manera que una característica detectada en qualsevol lloc es detecti a tot arreu— tampoc no es construeix aquí; és gairebé exactament correcte per a imatges i es delega a un curs de visió. Ni la recurrència ni la convolució tornen a aparèixer després d’aquesta pàgina, i per això cap de les dues té capítol: el capítol 1 prometia que les omissions es declararien, no que quedarien en silenci.

La funció més òbvia d’un nombre variable de vectors que retorna un vector és la mitjana:

ct=1ti=1txi\mathbf{c}_t = \frac{1}{t}\sum_{i=1}^{t} \mathbf{x}_i

Qualsevol nombre d’entrades, mida de sortida fixa, diferenciable, de franc. Una taula d’embedding més aquesta mitjana més una capa lineal fins al vocabulari és un model de llenguatge complet en quinze línies. També és terrible, i la manera com és terrible és tota la derivació.

El corpus següent és un megabyte de Shakespeare, 1.115.394 caràcters, passat per un tokenizer BPE a nivell de byte del tipus construït al capítol 7 amb un vocabulari de 1024: 459.760 tokens a 2,43 caràcters cadascun, dividit 90/10. Tots els models tenen amplada 128, veuen 128 tokens i s’entrenen durant 3000 passos d’AdamW a 10310^{-3} amb un batch de 64. La perplexitat és sobre la partició reservada.4

modelparàmetresperplexitat de validació
només el token actual, sense cap context263.16859,71
més la mitjana uniforme de tot el que hi ha abans263.168248,07
més embeddings de posició apresos279.552245,93
mitjana uniforme afegida al token en lloc de substituir-lo263.16860,45

Llegeix la segona fila dues vegades. Fer la mitjana del context no ajuda una mica; fa que el model sigui quatre vegades pitjor que ignorar el context del tot. Dues raons, totes dues demostrables més que empíriques.

La mitjana no pot veure l’ordre. La suma commuta, així que reordenar la finestra deixa el resum inalterat —no aproximadament:

order.pyPYTHON
A = torch.tril(torch.ones(T, T))
A = A / A.sum(1, keepdim=True)          # rows of the averaging matrix
y = x[torch.randperm(T)]                # the same tokens, shuffled
print((A[-1] @ x - A[-1] @ y).abs().max().item())
TEXT
2.9802322387695312e-08

Soroll de coma flotant en una suma reordenada: els dos resums són el mateix vector. Un model que només veu el context com una mitjana no pot distingir el gos va mossegar l’home de l’home va mossegar el gos. La tercera fila demostra que això no es pot arreglar afegint posicions a les entrades: un embedding de posició après en cada token abans de fer la mitjana va guanyar 2,14 punts de 188. Les posicions entren a la suma, i la suma les oblida.

I la mitjana ofega el present. A la posició 100, el token actual és una centèsima part del resum. Això té una reparació barata que ja tens: conserva el token i afegeix-hi el resum —una connexió residual, del capítol 6—, i la quarta fila mostra què fa. Un cop reparada la dilució, la mitjana uniforme no aporta res: 60,45 contra una línia base de 59,71. Tots els tokens hi són, ponderats igual, i ponderar igual és el mateix que no tenir informació.

El problema no és fer la mitjana. Són els pesos.

La mitjana és una multiplicació de matrius, i la màscara és una softmax

Enllaç a la secció: La mitjana és una multiplicació de matrius, i la màscara és una softmax

Fer la mitjana sobre un prefix creixent sembla un bucle. És una multiplicació per una matriu triangular inferior les files de la qual sumen u —i també, exactament, una softmax:

mechanics.pyPYTHON
loop = torch.stack([x[:t + 1].mean(0) for t in range(T)])   # the obvious version

A = torch.tril(torch.ones(T, T))
A = A / A.sum(1, keepdim=True)
mat = A @ x                                                  # the same thing

S = torch.zeros(T, T).masked_fill(torch.tril(torch.ones(T, T)) == 0, float("-inf"))
soft = F.softmax(S, dim=-1) @ x                              # and the same thing again
TEXT
loop vs matmul   max |diff| = 5.960464477539063e-08
loop vs softmax  max |diff| = 5.960464477539063e-08

the averaging matrix A (rows sum to 1, upper triangle is zero):
  1.000 0.000 0.000 0.000 0.000 0.000
  0.500 0.500 0.000 0.000 0.000 0.000
  0.333 0.333 0.333 0.000 0.000 0.000
  0.250 0.250 0.250 0.250 0.000 0.000
  0.200 0.200 0.200 0.200 0.200 0.000
  0.167 0.167 0.167 0.167 0.167 0.167

Tres components amb nom d’un transformer ara són a la pantalla. El triangle és la màscara causal, imposada per l’objectiu: si la posició tt pogués veure la posició t+1t{+}1, la resposta seria a l’entrada —la fuita que el capítol 6 et va dir d’auditar, però dins de l’arquitectura. La softmax és com s’implementa la màscara: posar les entrades prohibides a -\infty les envia exactament a zero i normalitza el que queda, de manera que emmascarar i normalitzar són una sola operació. (Fes servir -\infty, no -1e9: és el valor que la màscara vol dir, sobreviu a una conversió a float16 com -\infty, i t’estalvia decidir si la constant que has triat és prou gran per al rang on et trobes —que és la caixa de coma flotant del capítol 2 fent una pregunta que no cal respondre.) I les puntuacions són el paràmetre lliure. La mitjana uniforme és el que obtens quan cada puntuació permesa és el mateix nombre; posa-hi qualsevol nombre i la softmax el converteix en pesos vàlids.

La resta d’aquest capítol és una pregunta: d’on surten aquests nombres?

No poden ser paràmetres simples. Una matriu T×TT \times T apresa seria idèntica per a cada frase: podria codificar «mira quatre tokens enrere», però mai «mira el nom al qual fa referència aquest pronom». El pes que enllaça la posició tt amb la posició ii ha de dependre del que hi ha a totes dues posicions, perquè la rellevància és una relació, no una propietat: la paraula it no és intrínsecament rellevant, és rellevant per a alguna cosa.

La funció més barata de dos vectors que retorna un nombre és el producte escalar del capítol 1. Puntua la posició ii per a la posició tt com xtxi\mathbf{x}_t \cdot \mathbf{x}_i i el mecanisme funciona —malament, de dues maneres que forcen tota la resta. El producte escalar d’un vector amb si mateix és la seva norma al quadrat, així que cada token atendria sobretot a si mateix. I la relació seria simètrica: si it atén fortament a animal, llavors animal atén fortament a it, cosa falsa sobre el llenguatge, on un adjectiu necessita el seu nom molt més que el nom necessita l’adjectiu.

Així que dona a cada token dos rols, com dos mapes lineals apresos d’ell: allò que aquesta posició busca, qt=Wqxt\mathbf{q}_t = W_q\mathbf{x}_t, la query; i allò que ofereix perquè la trobin, ki=Wkxi\mathbf{k}_i = W_k\mathbf{x}_i, la key. Puntua qtki\mathbf{q}_t \cdot \mathbf{k}_i i la simetria desapareix, perquè WqWkW_q \neq W_k: un token pot anunciar una cosa i buscar-ne una altra.

Encara hi ha una cosa malament. La suma ponderada era sobre els xi\mathbf{x}_i mateixos, cosa que força que allò que es copia sigui allò que es fa coincidir. Fer coincidir vol les característiques que identifiquen un token; copiar vol les característiques útils més avall. Així que aprèn un tercer mapa, vi=Wvxi\mathbf{v}_i = W_v\mathbf{x}_i, el value, i suma aquests.

La fórmula ara és comptabilitat:

Attention(Q,K,V)=softmax ⁣(QKdk+M)V\mathrm{Attention}(Q, K, V) = \mathrm{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}} + M\right)V

amb MM com a màscara causal, zero a la diagonal i per sota, i -\infty per sobre. En codi són trenta línies, vint de les quals són formes:

attention.pyPYTHON
class Head(nn.Module):
    """One head of causal self-attention."""

    def __init__(self, d_model, d_head, block):
        super().__init__()
        self.q = nn.Linear(d_model, d_head, bias=False)      
        self.k = nn.Linear(d_model, d_head, bias=False)      
        self.v = nn.Linear(d_model, d_head, bias=False)      
        self.d_head = d_head
        self.register_buffer("mask", torch.tril(torch.ones(block, block)).bool())

    def forward(self, x):
        T = x.shape[1]
        q, k, v = self.q(x), self.k(x), self.v(x)
        s = q @ k.transpose(-2, -1) / math.sqrt(self.d_head)          
        s = s.masked_fill(~self.mask[:T, :T], float("-inf"))          
        w = F.softmax(s, dim=-1)                                      
        return w @ v                                                  

Puntua, emmascara, normalitza, barreja. Tota la resta és una projecció.

La divisió per l’arrel quadrada, i contra què defensa

Enllaç a la secció: La divisió per l’arrel quadrada, i contra què defensa

Gairebé totes les explicacions de dk\sqrt{d_k} diuen «per evitar que la softmax se saturi», cosa que és certa i no explica res. L’argument són dues línies de la variància del capítol 2. Si les entrades de q\mathbf{q} i k\mathbf{k} són independents amb mitjana zero i variància u, cada producte qjkjq_j k_j té variància u, i les variàncies de coses independents se sumen:

Var(qk)=j=1dkVar(qjkj)=dk\mathrm{Var}(\mathbf{q}\cdot\mathbf{k}) = \sum_{j=1}^{d_k}\mathrm{Var}(q_j k_j) = d_k

Així que les puntuacions tenen desviació estàndard dk\sqrt{d_k}. Mesurat sobre vint mil parelles aleatòries:

TEXT
     d     Var(q.k)         std   sqrt(d)
     4        3.975       1.994     2.000
    16       16.071       4.009     4.000
    64       64.249       8.016     8.000
   256      253.065      15.908    16.000
  1024     1015.562      31.868    32.000

Per què importa: la softmax és sensible a l’escala d’una manera que una capa lineal no ho és. Doblar l’entrada d’una capa lineal dobla la seva sortida; multiplicar les puntuacions per deu abans d’una softmax converteix una barreja suau en una tria dura. Una fila de 64 puntuacions, amb i sense la divisió:

dkd_kpes més gran, sense dividirentropiatokens efectiuspes més gran, dividitentropiatokens efectius
40,2052,94419,00,0813,75842,9
160,4381,6925,40,0753,84946,9
640,4890,8742,40,0853,67339,4
2560,99990,00071,00,1433,54734,7
10241,00000,00001,00,1323,64438,3

«Tokens efectius» és l’exponencial de l’entropia: sobre quantes posicions fa realment la mitjana la fila. Sense dividir, a dk=256d_k = 256, un head acabat d’inicialitzar atén exactament a un token de 64, triat només per la mostra aleatòria.

Això és dolent cap endavant i pitjor cap enrere, amb una forma que el capítol 5 ja va mesurar en una tanh\tanh. Una softmax compromesa amb una entrada gairebé no té derivada: la diagonal del seu jacobià és wi(1wi)w_i(1-w_i), zero als dos extrems. Sobre dues mil files aleatòries:

dkd_kiwi(1wi)\sum_i w_i(1-w_i) sense dividirdividitfiles saturades (pes més gran per sobre de 0,99)
40,84270,95680,2 % → 0,0 %
640,29400,960917,9 % → 0,0 %
2560,14060,960949,1 % → 0,0 %
10240,06810,961170,4 % → 0,0 %

A dk=1024d_k = 1024, set de cada deu files estan congelades abans que comenci l’entrenament, i un head que comença congelat no pot aprendre on mirar. Dividit, la quantitat és plana a 0,96 en cada amplada i res no se satura.

Ara la part que ningú no publica: canvia la perplexitat final? Elimina la divisió i entrena, amb quatre amplades de head:

amplada del headsense dividirdividit per dk\sqrt{d_k}dividit per dkd_k
quatre heads, dk=32d_k = 3237,2938,0737,89
un head, dk=128d_k = 12848,5146,1045,99
un head, dk=256d_k = 25665,3747,53
un head, dk=512d_k = 51267,0649,15
un head, dk=1024d_k = 102476,6959,17

Les dues primeres files venen del pressupost de 3000 passos anterior; les tres últimes són una execució més curta —1500 passos, batch de 32, un head, sense normalització abans de les projeccions— amb totes dues variants sota configuracions idèntiques.

A dk=32d_k = 32, la divisió no val res i l’execució sense ella va molt lleugerament per davant. Això no és una llicència per eliminar-la, perquè a 256 val 18 punts de perplexitat i a 1024 en val 17. El mecanisme és visible en les mateixes puntuacions:

dkd_kdesviació estàndard de la puntuació a l’inicidesprés de 1500 passos, sense dividirdesprés de 1500 passos, dividitfiles saturades, sense dividirdividit
25610,49121,672,1391,9 %0,8 %
51215,13836,852,6698,7 %1,3 %
102421,155147,463,4499,9 %16,5 %

El head sense dividir no es recupera. S’embala: la desviació estàndard de les seves puntuacions passa de 21 a la inicialització a 5147, l’entropia de l’attention cau a zero, i el 99,9 % de les files posen més de 0,99 del seu pes en un sol token. Un cop un head és un selector dur, el seu gradient és gairebé zero i res no el fa tornar, així que el col·lapse és estable. El head dividit queda en una desviació estàndard de puntuació de 3,44 després del mateix entrenament, que és una barreja suau que encara es pot canviar.

Vaswani et al. diuen exactament això i res més: sospiten que els productes «creixen molt en magnitud per a valors grans de dkd_k» i divideixen.5 La paraula gran porta pes, i les taules diuen on comença gran: res a 32, tot a 256.

Més d’una opinió, i els dos terços de què ningú parla

Enllaç a la secció: Més d’una opinió, i els dos terços de què ningú parla

Un head és una fila de softmax per posició, així que conté una resposta a «què és rellevant aquí». Predir la paraula després de the a the animal that crossed the wet street necessita alhora el lloc sintàctic, el subjecte i el token anterior, i una distribució de probabilitat no pot estar concentrada en tres llocs. Així que executa diversos heads en paral·lel, cadascun d’amplada dmodel/hd_{\text{model}}/h, concatena’ls i barreja’ls amb una matriu més WoW_o: has particionat l’amplada, no l’hi has afegit.

Attention també fa exactament una cosa: mou informació entre posicions. Cada operació del codi anterior és lineal al llarg de l’eix de característiques, i el capítol 5 va demostrar què és una pila de mapes lineals. Així que cada bloc també porta un petit MLP aplicat a cada posició independentment, que expandeix l’amplada per quatre i torna, amb un GELU al mig. Val la pena memoritzar la divisió de treball: attention barreja entre posicions, la xarxa feed-forward calcula dins d’una posició.

L’escala completa, cada fila afegint una peça a la fila de sobre:

modelparàmetresperplexitat de validació
mitjana uniforme, afegida279.55260,45
un head d’attention, substituint el token328.70455,47
un head d’attention, afegit328.70446,10
quatre heads en lloc d’un345.21643,21
més la xarxa feed-forward476.92839,87
més LayerNorm — el bloc complet477.69638,07

Els pesos apresos guanyen els uniformes per 14 punts de perplexitat, que és tot l’argument d’aquest capítol en una fila. Quatre heads en compren 3 més per 16.512 paràmetres extra. I el mateix head val 9 punts més afegit que substituint: attention porta informació cap endins, no decideix què és una posició.

Ara, on són realment els paràmetres, cosa que sorprèn qui només ha vist el diagrama:

ampladaheadsattentionfeed-forwardtotal per bloc
128465.664 (33,2 %)131.712 (66,6 %)197.888
768122.360.064 (33,3 %)4.722.432 (66,6 %)7.085.568
40963267.112.960 (33,3 %)134.238.208 (66,7 %)201.367.552

Dos terços de cada bloc transformer són la xarxa feed-forward, a qualsevol escala, perquè attention té quatre matrius d×dd \times d i l’MLP té l’equivalent de vuit. Sigui el que sigui que sap un model, la majoria dels paràmetres que ho sostenen són a l’MLP per posició.

Residuals i LayerNorm, heretats del capítol 6

Enllaç a la secció: Residuals i LayerNorm, heretats del capítol 6

LayerNorm es va construir i mesurar al capítol 6, i aquest capítol la fa servir tal com va quedar allà; les connexions residuals s’hi van anomenar i sotmetre a ablations, i aquí es construeixen. Les files «afegit, no substituint» de sobre són connexions residuals, que valen 188 punts de perplexitat per a la mitjana i 9 per a un head. LayerNorm7 normalitza cada exemple sobre les seves característiques, i el capítol 6 va donar les raons per les quals ella, i no BatchNorm, va sobreviure aquí —sense dependència del batch, sense estadístiques acumulades, idèntica en entrenament i inferència, indiferent a la longitud de la seqüència—, cadascuna de les quals esdevé un requisit quan generes un token cada vegada per a un usuari, que és on acaba el capítol 13. Costa 768 paràmetres i compra 1,8 punts de perplexitat.

block.pyPYTHON
class Block(nn.Module):
    def forward(self, x):
        x = x + self.att(self.ln1(x))     
        x = x + self.ff(self.ln2(x))      
        return x

Mira on seu la normalització: a l’entrada de cada subcapa, amb el camí residual de l’entrada a la sortida mai normalitzat. Això és pre-norm. L’article del 2017 fa el contrari, x = LayerNorm(x + Att(x))post-norm, que posa una LayerNorm al mateix camí residual.

Xiong et al. van explicar la diferència a través del gradient a la inicialització, que en una xarxa post-norm està mal escalat amb la profunditat —la raó per la qual el transformer original necessitava un warmup del learning rate per entrenar-se.8 Dotze blocs, 1000 passos, learning rate 3×1033 \times 10^{-3}:

TEXT
gradient norm per block at initialisation, before any step
  pre-norm    block 1 0.0498 ... block 12 0.0657   ratio last/first  1.32
  post-norm   block 1 0.0977 ... block 12 0.1613   ratio last/first  1.65

  pre-norm,  no warmup          perplexity   37.82
  pre-norm,  200-step warmup    perplexity   37.62
  post-norm, no warmup          perplexity  308.05
  post-norm, 200-step warmup    perplexity   37.88

Post-norm sense warmup és vuit vegades pitjor, i post-norm amb warmup iguala exactament pre-norm. Warmup no és una bona pràctica general aquí; és un pegat per a una disposició específica de la normalització, i moure la LayerNorm en treu la necessitat. Per això essencialment tots els models des del 2019 són pre-norm, i per això el diagrama del 2017 s’ha de llegir com a història més que no com a especificació.

Elimina els embeddings de posició i el model encara s’entrena; simplement no pot saber on és res, i això és una simetria més que no pas una fallada d’entrenament. Res en la puntuació d’attention esmenta tt o ii mateixos, així que permutar l’entrada permuta la sortida: self-attention és equivariant a permutacions. És la ceguesa a l’ordre de la mitjana amb una disfressa millor —la màscara causal restaura una mica d’ordre, perquè cada posició veu un prefix diferent, però dins d’un prefix tots els ordenaments són iguals.

Quatre maneres d’injectar posició, entrenades en finestres de 64 tokens i avaluades a 64, 128 i 256 —més enllà de qualsevol longitud que haguessin vist:

posicionsperplexitat a 64a 128a 256
cap48,7952,6357,52
embeddings absoluts apresos38,63108,47181,94
sinusoides fixes42,9695,26152,25
RoPE44,1250,5284,84
ALiBi44,9543,5142,49

Embeddings absoluts apresos —un vector per posició, afegit al token— guanyen a la longitud d’entrenament i després cauen pel precipici, perquè la posició 100 no havia estat mai en cap batch i el seu embedding encara és el vector aleatori amb què va començar. Sinusoides, l’elecció original, es calculen en lloc d’aprendre’s, a partir de sinus i cosinus a freqüències espaiades geomètricament; l’article del 2017 esperava que això extrapolés, i la taula diu que no: la funció està definida a la posició 200, però el model mai no va aprendre a llegir-la allà. RoPE9 no afegeix res i en canvi rota query i key per un angle proporcional a la posició, en talls bidimensionals; com que rotar igualment tots dos costats d’un producte escalar el deixa inalterat, la puntuació acaba depenent només de tit - i, així que la posició es torna relativa de franc i no hi ha cap taula que s’esgoti. Es degrada, però es degrada. ALiBi10 és el resultat més simple i més estrany aquí: una penalització lineal sobre la puntuació proporcional a la distància, amb un pendent diferent per head. La seva perplexitat millora a mesura que la finestra creix més enllà de la longitud d’entrenament, de 44,95 a 42,49, perquè la penalització està definida a qualsevol distància i cada head continua fent el que va ser entrenat per fer.

La lliçó sobreviu a la taula: una arquitectura que no pot representar alguna cosa és un problema diferent d’una que mai no va aprendre aquell rang, i el segon és el que mossega. També és la maquinària darrere de cada anunci de «hem ampliat el context a 128K»: gairebé sempre són reescalats d’una codificació rotatòria, i són la raó per la qual el capítol 16 diu que el límit de context es mou més que no pas desapareix.

Dropout s’hereta de la mateixa manera: apareix sobre els pesos d’attention després de la softmax, sobre la sortida de cada subcapa abans de l’addició residual, i sobre la suma d’embedding, fent exactament el que descrivia el capítol 6. En grans execucions de pretraining sovint es posa a zero, perquè un model que veu cada token una vegada no està en posició de sobreajustar.

Dos tensors de la capa tenen forma n×nn \times n, on nn és el nombre de tokens: les puntuacions i els pesos després de la softmax. Tota la resta —cada projecció, tot l’MLP— és lineal en nn.

Una capa d’attention, amplada 512, 8 heads, batch d’un, float32, en una GPU de portàtil. Llegeix les dues columnes de mil·lisegons només pels seus quocients: són temps de rellotge en una targeta de portàtil de 8 GB que baixa de 1.785 MHz a menys de 300 MHz quan s’escalfa, així que una execució en fred del mateix codi torna entre set i deu vegades més ràpida i una d’ocupada encara més lenta. Les columnes de megabytes són recomptes de bytes de l’assignador i no es mouen.

TEXT
  tokens   ms total    ms x4   ms projections   attn matrix MB    peak MB    MB x4
     128      2.246        -            1.324              0.5       14.6        -
     256      2.855     1.27            2.113              2.0       19.2     1.31
     512      5.761     2.02            3.105              8.0       34.4     1.79
    1024     16.414     2.85            4.008             32.0       89.1     2.59
    2048     51.573     3.14            9.989            128.0      296.1     3.32
    4096    225.432     4.37           20.176            512.0     1100.1     3.72
    8192    832.838     3.69           40.106           2048.0     4300.1     3.91
   16384   OUT OF MEMORY                                 8192.0

fitted exponent (log-log slope, last four rows):  time ~ n^1.91   memory ~ n^1.87

Les columnes x4 són el quocient respecte de la fila de sobre, i duplicar nn convergeix exactament a 4 tant per al temps com per a la memòria —3,91 a l’últim pas contra un 4 teòric. La columna de projeccions és el control: de 4,0 ms a 1024 tokens a 40,1 ms a 8192, un factor de deu per a un factor de vuit. Lineal, com s’havia promès.

Després, l’última fila. Una capa d’attention, una seqüència, sense cap model al voltant, es queda sense memòria en una GPU de 8 GB a 16.384 tokens: només la matriu de puntuacions seria de 8 GB, en ser 8 heads per 16.384 per 16.384 per 4 bytes. No el model; un tensor intermedi en una capa.

Aquest és el fet físic sota tres capítols posteriors. És per això que una context window té un límit, que el capítol 16 converteix en preu. És per això que existeix FlashAttention, que calcula el mateix resultat en rajoles sense emmagatzemar mai la matriu —una optimització de memòria abans que de velocitat.11 I és l’aritmètica darrere del preu d’un prompt llarg, que el capítol 24 paga en un bucle d’agent —un assumpte separat de l’altra troballa d’aquell capítol, que un model també utilitza pitjor un context llarg, cosa que mesura i es nega a atribuir a aquesta fórmula.

Mostra els detalls

Les dues variants que redueixen la cache, anomenades aquí i pagades al capítol 13.

La generació desa a la cache les keys i values dels tokens ja processats —una key i un value per token, per head per capa. Multi-query attention12 manté hh projeccions de query però una sola projecció de key i value compartida per tots els heads, dividint aquesta cache per hh. Grouped-query attention13 interpola: els heads s’agrupen, cada grup comparteix una key i un value, així que g=hg = h és attention ordinària i g=1g = 1 és multi-query. Gairebé tots els models oberts des del 2023 la fan servir amb 4 o 8 grups. Cap de les dues existeix per qualitat; totes dues existeixen per la mida d’aquesta cache, i el capítol 13 fa l’aritmètica que ho converteix en «quin model cap a la teva GPU».

L’article del 2017 descriu un encoder-decoder: una pila que llegeix la font amb attention sense màscara, una segona que genera l’objectiu causalment, i un tercer tipus d’attention al mig on les queries del decoder es troben amb les keys de l’encoder. Això és correcte per a traducció, on entrada i sortida són dues seqüències.

El que va guanyar va ser la meitat decoder-only: una pila, causal de cap a cap, entrada i sortida a la mateixa seqüència; i la raó no és l’elegància. «Predir el següent token» funciona sobre qualsevol text, així que el conjunt d’entrenament és internet en lloc d’un corpus paral·lel, i tot es converteix en aquesta única tasca: una traducció és un document que conté font i després objectiu, una pregunta i la seva resposta són un document, una conversa amb un tool call al mig és un document. El capítol 11 tracta de com es fabrica aquest últim. Els encoders no van desaparèixer: un veu tota l’entrada alhora, que és el que vols quan la feina és representar un text més que no pas continuar-lo, i per això els embeddings de recuperació del capítol 19 venen d’encoders i no del model que xateja.

Amb el bloc definit, la mida del model és aritmètica. Per bloc, amb amplada dd i una expansió de quatre vegades: 4d2+4d4d^2 + 4d per a Wq,Wk,Wv,WoW_q, W_k, W_v, W_o amb biaixos en totes quatre, com els té GPT-2 —la taula de sobre deixa el biaix fora de tres d’elles, d’aquí 2.304 menys per bloc a d=768d = 768; 8d2+5d8d^2 + 5d per a l’MLP; 4d4d per a dues LayerNorms— 12d2+13d12d^2 + 13d, més una taula de tokens de V×dV \times d i, per a posicions absolutes, nctx×dn_{\text{ctx}} \times d. Per a la forma de GPT-2 small —d=768d = 768, 12 blocs, un vocabulari de 50.257, un context de 1024, la capa de sortida compartint els pesos de l’embedding:

TEXT
  token embeddings     50,257 x 768 = 38,597,376
  position embeddings   1,024 x 768 =    786,432
  one block                             7,087,872
  12 blocks                            85,054,464
  final LayerNorm         2 x 768 =        1,536
  total (weights tied)                124,439,808

Que és la mida publicada d’aquest model. La fórmula no és una aproximació; és el model. Fixa’t també que gairebé un terç d’un model petit és la taula d’embedding, i per això la mida del vocabulari és una decisió arquitectònica i no de preprocessament —el trade-off que va preparar el capítol 7.

La perplexitat és un nombre sobre un corpus. El que fa un head és una pregunta diferent, i un model entrenat en un megabyte de Shakespeare és l’instrument equivocat per respondre-la: el més honest que es pot dir sobre el mapa d’attention d’un model de 500.000 paràmetres és que majoritàriament no és interpretable. Així que: un llenguatge on la pregunta té una resposta correcta.

La il·lustració clàssica és the animal did not cross the street because it was too tired, on it és l’animal, contra …because it was too wet, on una paraula mou el referent al carrer. Són esquemes de Winograd14: parelles de frases idèntiques excepte per una paraula, on aquesta paraula decideix a què fa referència un pronom.

També són resolubles fent trampa, que és la part que els tutorials ometen. Si els dos candidats són un animal i un lloc, tired i wet identifiquen el referent per categoria, i un model que només sap quines paraules hi són l’encerta sense saber res de l’ordre. Mesurat en aquesta versió de la tasca, amb parelles animal/lloc reservades:

TEXT
uniform causal average           held-out referent accuracy 100.0 %
one transformer block            held-out referent accuracy  91.7 %

La bossa de paraules guanya el transformer. Qualsevol demostració basada en aquesta frase no prova res sobre attention.

Així que tanca el forat: extreu tots dos candidats d’un sol conjunt de setze noms, qualsevol dels quals pot aparèixer en qualsevol ranura, i divideix els adjectius per rol en lloc de categoria —quatre que fan que it sigui qui creua (tired, scared, slow, weak), quatre que fan que sigui allò creuat (wet, wide, busy, steep).

TEXT
the {x} did not cross the {y} because it was too {adj} , so the {ref} waited .

Entrena com un predictor ordinari del següent token, puntua una posició —la paraula després de so the— i construeix el conjunt reservat a partir de parelles de noms l’ordre invertit de les quals era a l’entrenament, de manera que qualsevol cosa que sàpiga quins dos noms hi són però no quin va primer ha de respondre al revés.

modelparàmetresreservatanomena l’altre nom
només el token actual5.7965,2 %5,2 %
mitjana causal uniforme5.79627,9 %50,0 %
un head d’attention apresa18.08435,4 %64,6 %
quatre heads22.24475,0 %15,6 %
un bloc transformer55.71692,7 %4,2 %
dos blocs transformer105.508100,0 %0,0 %

L’atzar entre els dos noms presents és 50 %. La mitjana uniforme queda al 27,9 % i respon amb el nom equivocat de la parella exactament la meitat de les vegades: la signatura d’alguna cosa que sap quines paraules hi són i res sobre el seu ordre, tal com va predir la prova de reordenació tres seccions enrere.

Ara el mapa: l’attention a la posició que ha d’anomenar el referent, mitjanada sobre els quatre heads de cada bloc, per a les dues frases que difereixen en una paraula. Una mitjana uniforme posaria 0,067 en cadascun dels quinze tokens visibles.

TEXT
the animal did not cross the street because it was too tired , so the animal waited .
  blk 1  the:0.00 animal:0.70 did:0.00 not:0.00 cross:0.00 the:0.00 street:0.06
         because:0.00 it:0.00 was:0.00 too:0.00 tired:0.00 ,:0.05 so:0.00 the:0.19
  blk 2  the:0.00 animal:0.00 did:0.00 not:0.00 cross:0.00 the:0.00 street:0.00
         because:0.00 it:0.00 was:0.00 too:0.00 tired:1.00 ,:0.00 so:0.00 the:0.00

the animal did not cross the street because it was too wet , so the street waited .
  blk 1  the:0.00 animal:0.70 did:0.00 not:0.00 cross:0.00 the:0.00 street:0.06
         because:0.00 it:0.00 was:0.00 too:0.00   wet:0.00 ,:0.05 so:0.00 the:0.19
  blk 2  the:0.00 animal:0.00 did:0.00 not:0.00 cross:0.03 the:0.00 street:0.49
         because:0.00 it:0.00 was:0.00 too:0.20   wet:0.03 ,:0.00 so:0.00 the:0.25

El bloc 1 és idèntic en totes dues frases: 0,70 sobre el primer nom, sigui quin sigui l’adjectiu. Això no és un fracàs sinó una prova: a la primera capa, la query en una posició és una funció del token i l’índex propis d’aquella posició, i the a la posició 14 és el mateix token en totes dues frases. Un head de primera capa no pot condicionar-se a una paraula que encara no ha anat a buscar. Així que el bloc 1 fa l’única cosa útil disponible i arrossega el primer nom cap endavant.

El bloc 2 és on les frases se separen, i la mateixa fila sobre els vuit adjectius mostra la regla que el model va trobar:

adjectiubloc 2 sobre animalsobre streetsobre l’adjectiuresposta
tired, scared, slow, weak0,0000,0001,000animal
wet, wide, busy, steep0,0000,4910,00–0,03street

Per a un adjectiu de creuador, el segon bloc gasta tot el pes en l’adjectiu, perquè la resposta ja és al flux residual —el bloc 1 l’hi va posar— i tot el que necessita és confirmació. Per a un adjectiu de creuat, va a buscar l’altre nom. Això és un circuit de dos salts: un head mou un candidat endavant, un head d’una capa posterior llegeix un token que decideix si conservar-lo. La composició entre capes és el mecanisme, i és per això que un bloc va arribar al 92,7 % i dos van arribar al 100 %.

També és la forma del circuit més ben documentat en models reals. Induction heads —un head de token anterior que alimenta un head de la capa següent que completa el patró [A][B] … [A] → [B]— són el que el treball d’interpretabilitat d’Anthropic identifica darrere d’una gran part de l’in-context learning, i es formen en un moment identificable durant el pretraining. Aquest capítol no intenta aquesta anàlisi: es delega, amb tots dos articles a les referències, perquè llegir circuits d’un model real és un camp de recerca i no una secció.

Finalment, la implementació. Les trenta línies de dalt, amb els pesos copiats dels de PyTorch:

TEXT
ours vs nn.MultiheadAttention           max |diff| = 1.7881393432617188e-07
ours vs F.scaled_dot_product_attention  max |diff| = 1.7881393432617188e-07

1.8×1071.8 \times 10^{-7} en sortides amb magnitud mitjana 0,159: la mateixa aritmètica en un ordre diferent, amb precisió float32.

Tens l’arquitectura de què es construeix cada model de la resta d’aquest curs, i és més petita que la seva reputació: una mitjana ponderada amb pesos apresos, un MLP per posició que conté dos terços dels paràmetres, dues normalitzacions i dues addicions, apilades.

El que no tens és un model que sàpiga res, i apilar no ho arreglarà per si sol. Dos blocs en aquest corpus arriben a una perplexitat d’entrenament de 14,49 i una perplexitat de validació de 40,57, contra 18,77 i 38,07 d’un bloc —més capacitat, millor en el que ha vist, pitjor en el que no ha vist, que és la taula del capítol 6 amb un transformer dins. La distància entre aquest model i els que els capítols 14 a 30 interroguen no és arquitectònica. És el mateix bloc, repetit més vegades, sobre moltíssim més text.

Això ho converteix en un problema de comptabilitat, i la comptabilitat és més estranya del que sembla. Quant text, i d’on el treu ningú? Quanta aritmètica, i com l’estimes abans de gastar els diners? Donat un pressupost fix, és millor fer el model més gran o ensenyar-li més dades —i hi ha una resposta correcta, o només una moda? El capítol 10 respon a totes tres preguntes amb mesures, i posa preu a la forma útil més barata de la pregunta: què costa, avui, entrenar un model com GPT-2 des de zero?


Tres explicacions d’aquest material són millors que aquesta en allò per a què serveixen, i aquest capítol està escrit per llegir-lo al costat d’elles. The Illustrated Transformer de Jay Alammar és la millor imatge del flux de dades que s’ha dibuixat mai. The Annotated Transformer de Harvard NLP és l’article del 2017 amb codi en execució intercalat línia per línia. Let’s build GPT: from scratch, in code, spelled out d’Andrej Karpathy construeix el mateix model en directe en dues hores, i l’escala d’ablations de dalt és la mateixa columna vertebral mesurada en un corpus diferent. Per a la pregunta d’interpretabilitat que aquest capítol només toca, les fonts primàries són Elhage et al., A Mathematical Framework for Transformer Circuits (2021) i Olsson et al., In-context Learning and Induction Heads (2022), totes dues del grup d’interpretabilitat d’Anthropic.

  1. Hochreiter, S. and Schmidhuber, J. Long Short-Term Memory. Neural Computation 9(8), pp. 1735–1780 (1997).

  2. Sutskever, I., Vinyals, O. and Le, Q. V. Sequence to Sequence Learning with Neural Networks. arXiv:1409.3215 (2014). L’encoder-decoder en què l’únic vector de context és el coll d’ampolla.

  3. Bahdanau, D., Cho, K. and Bengio, Y. Neural Machine Translation by Jointly Learning to Align and Translate. arXiv:1409.0473 (2014). Attention, tres anys abans del transformer.

  4. La perplexitat és l’exponencial de l’entropia creuada mitjana per token, del capítol 8. Tots els nombres d’aquí fan servir el mateix tokenizer i la mateixa partició de validació, que és l’única condició sota la qual dues perplexitats es poden comparar.

  5. Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A. N., Kaiser, Ł. and Polosukhin, I. Attention Is All You Need. arXiv:1706.03762 (2017). La secció 3.2.1 és l’única frase sobre dk\sqrt{d_k} que aquest capítol dedica una secció a mesurar.

  6. Shazeer, N., Mirhoseini, A., Maziarz, K., Davis, A., Le, Q., Hinton, G. and Dean, J. Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer. arXiv:1701.06538 (2017).

  7. Ba, J. L., Kiros, J. R. and Hinton, G. E. Layer Normalization. arXiv:1607.06450 (2016). Introduïda i mesurada al capítol 6; usada aquí sense canvis.

  8. Xiong, R., Yang, Y., He, D., Zheng, K., Zheng, S., Xing, C., Zhang, H., Lan, Y., Wang, L. and Liu, T.-Y. On Layer Normalization in the Transformer Architecture. arXiv:2002.04745 (2020). L’anàlisi del gradient darrere de pre-norm, i l’argument que warmup és un símptoma.

  9. Su, J., Lu, Y., Pan, S., Murtadha, A., Wen, B. and Liu, Y. RoFormer: Enhanced Transformer with Rotary Position Embedding. arXiv:2104.09864 (2021).

  10. Press, O., Smith, N. A. and Lewis, M. Train Short, Test Long: Attention with Linear Biases Enables Input Length Extrapolation. arXiv:2108.12409 (2021). El resultat d’extrapolació reproduït més amunt.

  11. Dao, T., Fu, D. Y., Ermon, S., Rudra, A. and Ré, C. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. arXiv:2205.14135 (2022).

  12. Shazeer, N. Fast Transformer Decoding: One Write-Head is All You Need. arXiv:1911.02150 (2019).

  13. Ainslie, J., Lee-Thorp, J., de Jong, M., Zemlyanskiy, Y., Lebrón, F. and Sanghai, S. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. arXiv:2305.13245 (2023).

  14. Levesque, H. J., Davis, E. and Morgenstern, L. The Winograd Schema Challenge. KR (2012). La construcció darrere de la frase animal / street que fa servir cada tutorial d’attention.

A punt per deixar que triï LIA?

Crea amb tots els models d'IA en un sol lloc — comença gratis avui mateix.