Predicción del siguiente token: embeddings y qué significa la perplejidad
Entrena un modelo de caracteres con 32.033 nombres y observa cómo gradient descent redescubre una tabla de conteos a 4 decimales.
En esta página
Aquí tienes diez nombres producidos por un programa que nunca ha visto una palabra:
cexze momakurailezitynn konimittain llayn ka
da moliellavo emia sade ftlspNinguno de ellos es un nombre. Casi todos lo están intentando. Se pueden pronunciar, terminan donde terminan los nombres, y uno de ellos — emia — está a una sola letra de uno real. El programa que los produjo contiene 729 números, no tiene noción de palabra, sílaba ni persona, y se ajustó con una única pasada contando pares adyacentes de letras.
Al final de este capítulo, una red neuronal habrá reducido la puntuación de ese programa en un tercio sobre la misma medición. Lo que merece la pena ver es lo primero que hace la red: reproduce la tabla de conteos con tres decimales en cada fila bien poblada, sin que nadie se lo pida, porque ambos objetos son respuestas a la misma pregunta. Todo lo que viene después es lo que contar nunca podría haber hecho.
El objetivo es una identidad, no una decisión de diseño
Enlace a la sección: El objetivo es una identidad, no una decisión de diseñoEl capítulo 7 te dejó con una secuencia de enteros y ninguna razón para que uno siguiera a otro. Aquí está la razón, y es una línea del capítulo 2.
Un modelo de lenguaje es una función que toma los tokens vistos hasta ahora y devuelve una distribución sobre qué token viene después: un número por entrada del vocabulario, no negativo, y todos suman uno. Nada más. Para pasar de eso a una probabilidad de un documento entero, aplica la regla de la cadena de la probabilidad:
Eso es una identidad, verdadera para cualquier secuencia de cualquier cosa, sin supuestos añadidos. Así que un modelo que hace el trabajo pequeño — el siguiente token dados los anteriores — ya ha hecho el trabajo grande de asignar una probabilidad a cada documento posible, exactamente y gratis. La forma popular de presentarlo como un truco barato («solo predice la siguiente palabra») tiene la lógica al revés: predecir el siguiente token es modelar la distribución conjunta. Nunca hubo una segunda cosa que hacer.
La pérdida se deduce con la misma mecánica. En cada posición el modelo produce una distribución y la verdad es un único token conocido, así que la entropía cruzada del capítulo 4 se aplica sin cambios:
Esa es la log-verosimilitud negativa media: la receta del capítulo 2 con una distribución categórica en el hueco donde antes estaba la gaussiana. Y como la distribución verdadera es one-hot, su entropía es cero, así que por la identidad del capítulo 4 la entropía cruzada equivale a la divergencia KL: bajar este número y acercar las creencias del modelo a las de los datos son el mismo acto.
Una consecuencia merece su propia frase, porque es el hecho económico que sostiene todo el campo. Las etiquetas son los datos, desplazados una posición. Nadie anota nada. Un billón de tokens de texto son un billón de ejemplos ya etiquetados, por eso el corpus de entrenamiento de un modelo moderno es «internet» y no «un dataset que alguien construyó».
La línea base honesta: contar
Enlace a la sección: La línea base honesta: contarAntes de cualquier red, la línea base: 32.033 nombres, uno por línea, y la tarea de producir más nombres, letra a letra.1
El vocabulario son 26 letras más un símbolo de límite . que marca tanto el principio como el final de un nombre, así que el modelo tiene que aprender dónde empiezan los nombres y dónde se detienen. Son 27 símbolos, y el modelo más pequeño posible es una tabla de cuántas veces cada símbolo siguió a cada otro símbolo.
N = torch.zeros((27, 27), dtype=torch.int32)
for w in words:
cs = ["."] + list(w) + ["."]
for a, b in zip(cs, cs[1:]):
N[stoi[a], stoi[b]] += 1
P = N.float()
P = P / P.sum(1, keepdim=True) # one distribution per row Dos líneas de aritmética y el modelo queda ajustado; y no es una heurística: dividir los conteos por los totales de fila es la estimación de máxima verosimilitud para una distribución categórica, que es la receta del capítulo 2 con el cálculo ya hecho.
names: 32033 train/val/test: 25626 / 3203 / 3204
training bigrams: 182583
the six most likely letters after 'a':
a -> '.' 0.1944 a -> 'n' 0.1600 a -> 'r' 0.0967
a -> 'l' 0.0749 a -> 'h' 0.0690 a -> 'y' 0.0606Muestrea de él — elige una letra de la fila de la letra actual, pasa a esa fila, repite hasta que aparezca el símbolo de límite — y obtienes los nombres del principio de este capítulo. Fallan de una forma concreta e informativa: plausibles localmente, absurdos globalmente. Cada par adyacente de letras en momakurailezitynn es un par que aparece en nombres reales; simplemente hay diecisiete seguidos. El modelo tiene una letra de memoria, así que no puede saber que se está alargando demasiado.
Perplejidad, y cómo leerla
Enlace a la sección: Perplejidad, y cómo leerlaLa pérdida en nombres reservados es de 2,4546 nats. Ese número no significa nada por sí solo, por eso existe la perplejidad:
Escrito completo, sin que ninguna biblioteca haga el trabajo:
@torch.no_grad()
def perplexity(logits, Y):
logp = F.log_softmax(logits, dim=1) # log q for every symbol
chosen = logp[torch.arange(len(Y)), Y] # log q of the one that came next
return torch.exp(-chosen.mean()) Exponenciar deshace el logaritmo y devuelve el número a las unidades de contar cosas. La forma clara de ver qué cuenta es medir un modelo que no sabe absolutamente nada: uno que asigna probabilidad a cada símbolo con independencia del contexto:
uniform over 27 symbols loss 3.2958 nats ppl 27.000
bigram counts, add-one smoothed loss 2.4546 nats ppl 11.642Exactamente 27,000, porque . La perplejidad es el número efectivo de opciones igualmente probables entre las que el modelo está eligiendo. Una perplejidad de 27 significa «ni idea, podría ser cualquier cosa». El 11,642 del modelo de conteos significa que una letra de contexto lo deja tan incierto como alguien que elige a ciegas entre unas doce opciones en lugar de veintisiete; por eso se cita la perplejidad y no la pérdida bruta.
Dos cosas salen mal con ella, y la segunda sale mal en artículos publicados.
Las probabilidades cero son fatales. De las 729 celdas de la tabla, 113 nunca aparecen en entrenamiento: el 15,5 % está vacío. Eso está bien hasta que el conjunto reservado cae en una, y siete bigramas de validación lo hacen, entre ellos d→q, z→j y q→o dos veces. Probabilidad cero significa log , lo que significa pérdida infinita y perplejidad infinita: un nombre entre tres mil destruye la métrica. El parche habitual es sumar 1 a cada conteo antes de normalizar, lo que aquí cuesta casi nada (2,4546 en lugar de 2,4524). Pero el parche es una confesión. Un modelo de conteos no puede generalizar en absoluto. No tiene forma de sospechar que q→o es plausible porque q→u es común y o se comporta como u en otros lugares, ya que no tiene noción de que dos símbolos puedan parecerse. Cada celda se aprende por separado, y arreglar eso es el objetivo del resto de este capítulo.
La perplejidad es un precio por token, y el token es un parámetro libre. Este es el error que aparece constantemente cuando se comparan modelos, y es fácil verlo en cuanto miras. Toma el mismo corpus de prosa inglesa del capítulo 7, el mismo modelo de bigramas interpolado, y cambia solo cómo se trocea el texto:
| unidad | vocabulario | tokens en test | entropía cruzada | perplejidad | bits por carácter |
|---|---|---|---|---|---|
| caracteres | 76 | 14.469 | 2,5217 | 12,45 | 3,6378 |
| BPE, 512 fusiones | 329 | 6.871 | 3,8547 | 47,21 | 2,6407 |
| BPE, 2.048 fusiones | 1.820 | 4.233 | 5,7468 | 313,20 | 2,4254 |
| palabras | 2.991 | 6.284 | 3,5627 | 35,26 | 2,2322 |
La perplejidad varía por un factor de 25 entre esas filas. Nada del modelo cambió; solo el tamaño de lo que se predice. Predecir una palabra entera es más difícil que predecir una letra, así que cuesta más por predicción, y hay menos predicciones que hacer.
Ahora lee la última columna, que divide el coste total por el número de caracteres y lo convierte a bits. Reordena la tabla. Por perplejidad, el ranking es caracteres, palabras, BPE-512, BPE-2048; por bits por carácter es palabras, BPE-2048, BPE-512, caracteres. El modelo de caracteres pasa del primer puesto al último. El modelo de 2.048 fusiones, que por perplejidad parece 6,6 veces peor que el de 512 fusiones, es en realidad el mejor de los dos: 2,4254 bits frente a 2,6407.
Así que una perplejidad solo es comparable entre dos modelos que comparten tokenizer, y los modelos con tokenizers distintos solo pueden compararse en bits por carácter: la cantidad que Shannon midió en 1951 haciendo que sujetos humanos adivinaran la siguiente letra de texto inglés, y que acotó en torno a un bit por carácter.2 Nuestro mejor bigrama está en 2,23 bits, lo que resume bastante bien cuánto le queda por recorrer a este capítulo.
Lo mismo, aprendido
Enlace a la sección: Lo mismo, aprendidoAhora construye el mismo modelo como una red. Necesitará órdenes de magnitud más aritmética para llegar al mismo sitio, y llegar al mismo sitio es el punto.
Sustituye la tabla por una matriz de pesos de forma . Convierte la letra actual en un vector one-hot, multiplica, y llama logits al resultado: las puntuaciones sin normalizar del capítulo 4. Luego softmax, luego entropía cruzada, luego gradient descent.
W = torch.randn((27, 27), requires_grad=True)
for step in range(3000):
logits = W[xs]
loss = F.cross_entropy(logits, ys)
W.grad = None
loss.backward()
W.data -= 50.0 * W.gradLa línea destacada contiene una definición que conviene tener. Multiplicar un vector one-hot por una matriz selecciona una fila de ella, así que la multiplicación es una búsqueda; y toda implementación se salta la aritmética y hace la búsqueda directamente, que es lo que es W[xs].
Eso es una embedding table. Una matriz con una fila por entrada del vocabulario, indexada por token id. Sin geometría, sin semántica, sin algoritmo aparte: una tabla de búsqueda cuyo contenido resulta que se aprende por gradient descent junto con todo lo demás. Toda afirmación mística sobre el «embedding space» termina aquí.
Entrénala y observa adónde va:
step 1 train 3.7550 val 3.3882 max gap to the count table 0.757269
step 100 train 2.4732 val 2.4726 max gap to the count table 0.388354
step 1000 train 2.4557 val 2.4549 max gap to the count table 0.041862
step 3000 train 2.4547 val 2.4544 max gap to the count table 0.004048La última columna es la mayor diferencia absoluta entre cualquier celda de softmax(W) y la celda correspondiente de la tabla de conteos, y tiende a cero. Tras 3.000 pasos, el mayor desacuerdo en cualquiera de las 729 celdas es 0,004048 y la media es 0,000224. La peor celda es q→i, vista doce veces en todo el conjunto de entrenamiento; entre las 22 filas con más de mil apariciones, el peor desacuerdo es 0,000562.
count table network
a -> '.' 0.1945 0.1945
a -> 'n' 0.1601 0.1601
a -> 'r' 0.0967 0.0967Gradient descent, partiendo de números aleatorios y sin más instrucción que «haz grande la log-probabilidad de la siguiente letra», redescubrió la tabla de conteos. Y tenía que hacerlo: los conteos son la estimación de máxima verosimilitud, la entropía cruzada es la log-verosimilitud negativa, así que ambos procedimientos optimizan el mismo objetivo y ese objetivo tiene un óptimo. La red no aprendió algo parecido a contar. Convergió a contar, despacio.
Lo que plantea la pregunta justa de por qué alguien se molestaría. Porque la tabla de conteos no tiene adónde ir desde aquí, y la red sí.
El contexto es el cuello de botella, no la capacidad
Enlace a la sección: El contexto es el cuello de botella, no la capacidadExtiende el modelo para que mire más de un carácter anterior. Esta es la arquitectura de Bengio de 2003, el ancestro directo de todos los modelos del resto de este curso:4 toma los tres últimos caracteres, mapea cada uno mediante una embedding table a una fila de 10 dimensiones, concatena las filas en 30 números, pásalos por la capa oculta del capítulo 5, y termina con una capa de salida que produce un logit por entrada del vocabulario.
C = torch.randn((27, 10)) # the embedding table
W1 = torch.randn((3 * 10, 200)) # the hidden layer from Chapter 5
W2 = torch.randn((200, 27)) # one output per vocabulary entry
emb = C[X].view(-1, 30) # three lookups, concatenated
h = torch.tanh(emb @ W1 + b1)
logits = h @ W2 + b2
loss = F.cross_entropy(logits, Y)Fíjate en qué es nuevo y qué no. La capa oculta es la del capítulo 5, sin cambios; la pérdida es la del capítulo 4, sin cambios. Las novedades son la embedding table al principio y una capa de salida tan ancha como el vocabulario del capítulo 7; y esta segunda es la parte cara de todo modelo de lenguaje que se haya construido jamás, porque un vocabulario real tiene 100.000 entradas y esta multiplicación de matrices se ejecuta en cada posición.
El mismo código, entrenado de forma idéntica, cambiando solo el tamaño de la context window:
| contexto | parámetros | pérdida de validación | perplejidad de validación |
|---|---|---|---|
| conteo, 1 carácter | 729 | 2,4546 | 11,642 |
| neuronal, 1 carácter | 7.897 | 2,4577 | 11,678 |
| neuronal, 3 caracteres | 11.897 | 2,1145 | 8,285 |
| neuronal, 8 caracteres | 21.897 | 2,0506 | 7,773 |
La segunda fila es la interesante. Una red con una capa oculta de 200 unidades y once veces más parámetros que la tabla de conteos rinde exactamente igual que la tabla de conteos y nada mejor. La capacidad nunca fue la limitación. Un carácter de contexto permite cierta pérdida y nada que le añadas puede bajar de ahí, porque la información no está.
Dale tres caracteres y la perplejidad baja de 11,68 a 8,29: una reducción del 29 %, comprada con 4.000 parámetros extra. Supera al conteo aquí precisamente por la razón diagnosticada antes: un modelo de conteos sobre contextos de tres caracteres necesita filas, la mayoría vacías o con una sola observación, y aprende cada una por separado. La red comparte. Si a, e y i acaban con filas de embedding similares, lo que aprende después de bra se transfiere a bre sin haber visto nunca bre. Esa transferencia es todo el valor de la embedding table, y es la brecha entre las filas dos y tres.
Las muestras mejoran en consecuencia:
deliah nellara joce kael quintis
salayson reety khyrmin mahnen madiaryxiaSigue sin ser una lista de nombres reales. Pero deliah, nellara y kael no desentonarían en una, y los monstruos interminables han desaparecido: el más largo de veinte ejemplos del modelo de conteos tiene diecinueve letras; el más largo de veinte de este tiene trece.
Qué hay realmente dentro de la embedding table
Enlace a la sección: Qué hay realmente dentro de la embedding tableLa tabla es : una fila de diez números por carácter, todos inicializados aleatoriamente y movidos solo por el gradiente de la pérdida del siguiente carácter. Nadie puso nada ahí. Entonces, ¿qué acabó dentro?
La herramienta para preguntarlo es la similitud del coseno, que es el producto escalar del capítulo 1 con las longitudes divididas:
Mide el ángulo entre dos vectores e ignora sus longitudes, que es lo que quieres cuando la longitud de una fila refleja cuántas veces apareció su token y no lo que significa. Normaliza primero cada vector a longitud 1 — como hacen los sistemas reales, una vez, en el momento de indexar — y la similitud del coseno es simplemente el producto escalar.
Estos son los vecinos más cercanos de algunos caracteres en la tabla entrenada:
'c' -> 'k':+0.598 'j' -> 'z':+0.650 'i' -> 'y':+0.541
'u' -> 'e':+0.482 'a' -> 'h':+0.367 '.' -> 'q':+0.077Parte de eso es lo que promete el folklore. c y k son intercambiables en nombres, igual que i y y; j y z son consonantes raras, casi siempre iniciales, que se comportan de forma parecida. El símbolo de límite . no está cerca de nada: 0,077 con la letra más próxima, porque es el único símbolo que marca una posición en lugar de un sonido.
Y parte no. El vecino más cercano de a es h, no otra vocal. Promediado sobre todos los pares:
mean cosine, vowel to vowel : +0.1889
mean cosine, consonant to consonant : +0.0765
mean cosine, vowel to consonant : -0.0042Las vocales se parecen más entre sí que a las consonantes, y el efecto es real pero pequeño. Probado contra 2.000 grupos de cinco letras elegidos al azar, 58 de esos grupos se separan al menos igual de limpiamente: una brecha significativa en torno a . Real, entonces, pero nada parecido a la isla geométrica nítida que sugieren las explicaciones populares de embeddings.
Esa es la descripción honesta de una embedding table y merece la pena conservarla durante el resto del curso. No es un mapa de significado. Es un cambio de coordenadas, aprendido en lugar de diseñado, cuya única tarea es facilitar el trabajo de la siguiente capa: la misma frase que el capítulo 5 usó para la capa oculta que plegaba el plano para resolver XOR. Cualquier estructura que encuentres en ella está ahí porque redujo la pérdida, y la estructura que no reduce la pérdida simplemente no está.
word2vec, GloVe y la aritmética que todo el mundo cita
Enlace a la sección: word2vec, GloVe y la aritmética que todo el mundo citaSi la parte útil es la tabla, puedes ir directamente a por ella. Eso es word2vec: conserva la búsqueda de embedding, tira el modelo de lenguaje.5
El objetivo de skip-gram with negative sampling cabe en una línea. Para un par real (centro, contexto) extraído del corpus, empuja hacia arriba su producto escalar; para pares falsos extraídos de una distribución de ruido, empújalo hacia abajo:6
Eso es una clasificación binaria — «¿aparecieron realmente juntas estas dos palabras?» — y es barato precisamente porque nunca toca el vocabulario completo, que es lo que hizo práctico entrenar con miles de millones de palabras en 2013. GloVe llega a vectores similares desde la otra dirección, factorizando la matriz de conteos globales de coocurrencia en lugar de recorrer ejemplos en streaming.7 Ambos se ajustan exactamente a la estadística de la que se construyó la tabla de conteos. Son conteo, comprimido.
Entrenados en text8 — 17.005.207 palabras de la Wikipedia en inglés, 71.290 de ellas con al menos cinco apariciones, 100 dimensiones, tres pasadas — los vectores salen con la propiedad que los hizo famosos:
king -> charles 0.700, son 0.693, queen 0.686, henry 0.669, throne 0.667
physics -> chemistry 0.672, electromagnetism 0.661, quantum 0.654, theoretical 0.624
guitar -> bass 0.733, vocals 0.732, acoustic 0.728, guitars 0.703, drums 0.685
three -> seven 0.892, two 0.877, one 0.875, five 0.871, four 0.870Nadie proporcionó una categoría para instrumentos ni para numerales. Ahora la parte famosa: toma king, resta man, suma woman y encuentra el vector más cercano al resultado.
king - man + woman
nothing excluded : king 0.693, elizabeth 0.657, wife 0.629, woman 0.607
a, b, c excluded : elizabeth 0.657, wife 0.629, mary 0.607 (queen is 4th, 0.604)El vector más cercano a king - man + woman es king. No es una rareza de un ejemplo. El conjunto de evaluación de Mikolov plantea preguntas de la forma a : b :: c : ? — 8.869 semánticas (paris : france :: rome : italy) y 10.675 sintácticas (walking : walked :: swimming : swam) — y, entre las 4.103 preguntas semánticas que este vocabulario puede responder, la ganadora es una de las tres palabras de entrada el 99,8 % de las veces. Las demostraciones publicadas no lo mencionan, porque la regla de puntuación estándar elimina a, b y c antes de mirar. Es una regla legítima, y está haciendo más trabajo que la aritmética:
| cómo se elige la respuesta | semánticas | sintácticas |
|---|---|---|
| offset, con las entradas excluidas (estándar) | 17,0 % | 11,9 % |
| offset, sin excluir nada | 0,1 % | 0,4 % |
vecino más cercano solo de c, entradas excluidas | 13,1 % | 9,3 % |
vecino más cercano solo de b, entradas excluidas | 2,3 % | 0,4 % |
La tercera fila es la que conviene digerir. Tira a y b, no hagas ninguna aritmética, devuelve lo que esté más cerca de c, y conservas el 77 % de la puntuación semántica. La mayor parte de lo que parece razonamiento analógico es proximidad más una regla que prohíbe las respuestas obvias, que es lo que Linzen midió con vectores entrenados correctamente y lo que las líneas base anteriores replican.8 Estos vectores concretos son pequeños — 17 millones de palabras frente a los miles de millones detrás de los modelos publicados — así que lee los porcentajes como una forma, no como el estado del arte. La forma es lo que sobrevive a cualquier escala: la aritmética es real, y mucho más débil que la demostración que todo el mundo cita.
Estático y contextual: un vector por palabra, o uno por aparición
Enlace a la sección: Estático y contextual: un vector por palabra, o uno por apariciónTodo lo visto hasta ahora tiene un límite rígido incorporado en la estructura de datos. Una tabla tiene una fila por token. La palabra bank recibe un vector, el mismo en una frase sobre un río y en una frase sobre una hipoteca; necesariamente, porque una búsqueda por id no puede depender de nada más.
La solución es dejar de leer el vector de la tabla y empezar a calcularlo a partir de la frase. Eso es un contextual embedding, introducido por ELMo en 2018 y convertido en estándar por BERT ese mismo año.910 Medidos en el modelo real, los números son más claros que la explicación:
sentence A: "He sat on the bank of the river and watched the water go by."
sentence B: "She deposited the cheque at the bank on the corner of the street."
static vector for 'bank' (a row of the input embedding table)
cosine A vs B ........................ 1.000000
contextual vector for 'bank', layer by layer
layer | A vs B | A vs another river sentence | B vs another money sentence
0 | 0.9512 | 0.9512 | 0.9359
4 | 0.5647 | 0.8987 | 0.7716
9 | 0.4284 | 0.8699 | 0.7568
12 | 0.5278 | 0.8702 | 0.7335La primera fila es exacta, no aproximada: el vector estático de bank son los mismos 768 números en ambas frases, así que el coseno es 1 por construcción. Nueve capas después, las dos apariciones están en 0,43, mientras que bank en dos frases distintas sobre ríos se mantiene en 0,87. Nadie etiquetó ningún sentido en este proceso; los sentidos se separaron porque separarlos hace que el objetivo de entrenamiento — adivinar un token oculto a partir de sus vecinos — sea más fácil de satisfacer.
Dos detalles merecen atención. La capa 0 ya está en 0,9512 en lugar de 1,0, porque se han añadido position embeddings y la palabra ocupa un lugar distinto en cada frase. Y la similitud vuelve a subir en las capas 11 y 12: las capas finales de un modelo preentrenado están especializadas en su objetivo de entrenamiento, y a menudo no son el mejor sitio del que tomar una representación.
Mostrar detalles
Opcional: weight tying.
En bert-base-uncased la embedding table es : 23.440.896 números, el 21,4 % de los 109.482.240 parámetros del modelo. En un modelo de lenguaje pequeño la fracción es aún mayor, por eso un truco es casi universal: la tabla de entrada y la capa de salida que produce los logits son la misma matriz, usada una vez mediante búsqueda de filas y otra transpuesta.11 La capa de salida ya asigna un vector a cada entrada del vocabulario — toma un producto escalar contra cada una — y tying dice que el vector usado para leer un token y el vector usado para escribirlo deberían ser el mismo objeto. Recorta parámetros y mejora la perplejidad a la vez, algo lo bastante raro como para fijarse.
Un embedding model no es un modelo de lenguaje
Enlace a la sección: Un embedding model no es un modelo de lenguajePara buscar en un corpus por significado necesitas un vector por frase. Con eso, la búsqueda es trivial: esto es todo el retrieval semántico, y el capítulo 19 trata de todo lo que hay alrededor:
E = normalise(embed(sentences)) # (200, d), every row of length 1
q = normalise(embed([query])) # (1, d)
scores = q @ E.T # one matrix multiply
top5 = scores[0].argsort()[::-1][:5]Así que la única pregunta real es de dónde sale embed. El movimiento obvio es tomar un modelo de lenguaje preentrenado, pasar cada frase por él y promediar los vectores de token. Aquí está ese método frente a cuatro alternativas, puntuado de dos formas: la correlación de rangos entre coseno y juicios humanos de similitud sobre los 1.379 pares del benchmark STS, y retrieval top-1 en un índice construido con los 200 pares más fuertemente parafraseados de esos pares: un lado de cada par indexado, el otro usado como consulta.
| cómo se embebe la frase | correlación de rangos | top-1 en un índice de 200 frases |
|---|---|---|
| solapamiento binario de palabras (sin modelo) | 0,5500 | 89,0 % |
| media de los vectores estáticos entrenados arriba | 0,5263 | 85,5 % |
BERT, el token [CLS] | 0,2030 | 67,0 % |
| BERT, media de vectores de token | 0,4729 | 84,0 % |
| MiniLM, entrenado contrastivamente | 0,8203 | 92,0 % |
Lee las tres filas centrales contra las dos primeras. Un transformer preentrenado de 109 millones de parámetros, usado de la forma obvia, es peor juzgando similitud entre frases que contar cuántas palabras comparten dos frases; y peor que promediar los vectores text8 de 100 dimensiones entrenados hace un momento. El token [CLS], que muchos tutoriales aún recomiendan porque BERT se preentrenó con un objetivo a nivel de frase asociado a él, es peor que la mitad de eso.
Esto no es un defecto de BERT. Es el objetivo. Un modelo de lenguaje se entrena para que sus estados ocultos predigan un token; nada ahí pide que dos paráfrasis acaben cerca, y nada recompensa una geometría en la que el coseno signifique «mismo significado». La última fila es un modelo de una quinta parte del tamaño (22.713.216 parámetros) entrenado con una pérdida completamente distinta: aprendizaje contrastivo, donde los ejemplos son pares — una pregunta y su respuesta, una frase y su paráfrasis — y el objetivo acerca los pares verdaderos mientras aleja negativos muestreados. Esa es la contribución de Sentence-BERT y el origen de toda la industria de los embedding models.12 Dense Passage Retrieval aplica la misma receta directamente a la búsqueda, con un encoder para consultas y otro para pasajes.13
Así que, la regla práctica:
Un embedding model no es un modelo de lenguaje con la última capa eliminada. Es un modelo distinto con un objetivo distinto, normalmente mucho más pequeño, cuyo coseno significa lo que quieres que signifique porque se entrenó con pares donde ese era el objetivo. La tabla anterior es el coste de sustituir uno por otro.
Y la familia falla con el orden de las palabras. «The dog bit the man» y «the man bit the dog» tienen bolsas de palabras idénticas, así que el solapamiento de palabras y la media de vectores estáticos les dan un coseno exactamente 1,000000, y BERT con mean pooling, que sí ve la posición, aun así cae casi ahí; e incluso MiniLM entrenado contrastivamente los pone en 0,979. Si tu tarea de retrieval depende de quién hizo qué a quién, ningún umbral de coseno te salvará.
El capítulo 19 construye un sistema de retrieval de producción sobre esta base y llega a un corte de coseno concreto. La última medición de este capítulo es lo que hace que un número así sea defendible en lugar de mágico.
La maldición de la dimensionalidad, en una tabla
Enlace a la sección: La maldición de la dimensionalidad, en una tablaLos embeddings reales tienen cientos o miles de componentes, y las distancias se comportan de forma extraña ahí arriba. Toma 1.000 puntos aleatorios en el cubo unidad de dimensiones y mira la razón entre la distancia mayor y la menor entre dos cualesquiera de ellos:
| dimensiones | par más cercano | par más lejano | razón |
|---|---|---|---|
| 2 | 0,0007 | 1,3612 | 1921,66 |
| 10 | 0,2361 | 2,3397 | 9,91 |
| 100 | 3,0047 | 5,1752 | 1,72 |
| 1.000 | 11,7809 | 14,0306 | 1,19 |
| 10.000 | 39,6152 | 42,0125 | 1,06 |
En diez mil dimensiones, el par de puntos más lejano está solo un 6 % más separado que el par más cercano. Todo está aproximadamente a la misma distancia de todo lo demás, «vecino más cercano» deja de llevar mucha información, y eso es la maldición de la dimensionalidad, además de una razón por la que las grandes bases de datos vectoriales no hacen búsqueda exacta de vecinos más cercanos. La otra cara de la misma moneda es lo que hace viables los umbrales de coseno: medido sobre mil pares de vectores unitarios aleatorios, el coseno medio está en en 100 dimensiones y en en 768, con desviaciones estándar de 0,0968 y 0,0357; y en 768 dimensiones solo el 0,2 % de los pares aleatorios supera 0,1 en valor absoluto. Por tanto, una similitud medida de 0,4 no significa «se parecen en un 40 %»; está muy fuera de cualquier cosa que produzca el azar, por eso umbrales entre 0,3 y 0,7 separan señal de ruido en lugar de situarse en medio.
Adónde va esto ahora
Enlace a la sección: Adónde va esto ahoraEl modelo de este capítulo lee un número fijo de caracteres anteriores, busca cada uno y pega los resultados en orden. Ese diseño tiene dos problemas, y son el mismo problema.
Vuelve a mirar la tabla de contexto: pasar de tres caracteres a ocho casi duplicó los parámetros y compró 0,06 nats. El coste crece linealmente con el contexto — cada posición extra necesita su propia losa de la primera matriz de pesos — y el beneficio no. Llévalo a mil tokens y la primera capa por sí sola pesa más que el resto del modelo, con la mayor parte gastada en posiciones que no importan para una predicción dada.
Ese es el segundo problema: el modelo no tiene forma de decidir cuáles de los tokens anteriores importan. La posición dos tiene sus propios pesos y la posición siete los suyos, permanentemente, sea lo que sea que haya en ellas. Cuando el modelo está deletreando nell, el carácter decisivo es el inmediatamente anterior. Cuando una frase contiene un pronombre, la palabra que fija su referente puede estar cuarenta tokens atrás; y no se puede asignar una ranura fija a «cuarenta atrás», porque la próxima vez serán seis.
Lo que queremos es un modelo que calcule, para cada predicción, cuánto debe contar cada token anterior: pesos sobre el contexto producidos por el contenido en lugar de fijados por la disposición. Escríbelo con cuidado y empieza como algo completamente mundano: una media sobre los tokens anteriores. Luego deja que los pesos de esa media se aprendan, y deja que dependan de qué token está haciendo la pregunta.
Eso es attention, y es el capítulo 9.
Fuentes y método
Enlace a la sección: Fuentes y métodoTambién merece la pena leer en paralelo: el capítulo 3 de Speech and Language Processing, de Jurafsky y Martin, que trata los modelos n-gram, el smoothing y la perplejidad con mucho más cuidado del que cabe aquí, incluyendo por qué la interpolación y el back-off superan a sumar uno; las notas de Stanford CS229 §17.1–17.2 para el modelado de lenguaje desde el lado probabilístico; y el artículo de Linzen citado arriba, que es breve y vale la pena leer entero.
Referencias
Enlace a la sección: Referencias-
El ejemplo de generación de nombres, el dataset y la progresión desde una tabla de conteos hasta una red al estilo Bengio siguen la serie building makemore de Andrej Karpathy, cuyas dos primeras partes son el mejor acompañamiento para este capítulo. ↩
-
Shannon, C. E. Prediction and Entropy of Printed English. Bell System Technical Journal 30(1), pp. 50–64 (1951). Sujetos humanos adivinando la siguiente letra de texto inglés, y la medición original de bits por carácter. ↩
-
Shannon, C. E. A Mathematical Theory of Communication. Bell System Technical Journal 27 (1948). El teorema de codificación de fuente y la identificación de predicción con compresión. ↩
-
Bengio, Y., Ducharme, R., Vincent, P. and Jauvin, C. A Neural Probabilistic Language Model. Journal of Machine Learning Research 3, pp. 1137–1155 (2003). La arquitectura usada arriba: un embedding por palabra, concatenado sobre una ventana fija, a través de una capa oculta, hasta una softmax sobre el vocabulario. ↩
-
Mikolov, T., Chen, K., Corrado, G. and Dean, J. Efficient Estimation of Word Representations in Vector Space. arXiv:1301.3781 (2013). CBOW y skip-gram, y el conjunto de analogías usado arriba. ↩
-
Mikolov, T., Sutskever, I., Chen, K., Corrado, G. and Dean, J. Distributed Representations of Words and Phrases and their Compositionality. arXiv:1310.4546 (2013). Negative sampling, submuestreo de palabras frecuentes y la distribución de ruido elevada a la potencia 3/4 usada arriba. ↩
-
Pennington, J., Socher, R. and Manning, C. GloVe: Global Vectors for Word Representation. EMNLP 2014. Vectores de palabras a partir de una factorización de la matriz global de coocurrencia en lugar de ventanas locales recorridas en streaming. ↩
-
Linzen, T. Issues in evaluating semantic spaces using word analogies. RepEval 2016, arXiv:1606.07736. La fuente de las líneas base sin offset replicadas arriba. ↩
-
Peters, M. et al. Deep contextualized word representations. arXiv:1802.05365 (2018). ELMo: un vector por aparición, calculado por un modelo de lenguaje bidireccional. ↩
-
Devlin, J., Chang, M.-W., Lee, K. and Toutanova, K. BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. arXiv:1810.04805 (2018). El modelo medido en el experimento de bank. ↩
-
Press, O. and Wolf, L. Using the Output Embedding to Improve Language Models. arXiv:1608.05859 (2016), e Inan, H., Khosravi, K. and Socher, R. Tying Word Vectors and Word Classifiers. arXiv:1611.01462 (2016). Dos argumentos independientes para el mismo truco. ↩
-
Reimers, N. and Gurevych, I. Sentence-BERT: Sentence Embeddings using Siamese BERT-Networks. arXiv:1908.10084 (2019). Su medición inicial — BERT con mean pooling rindiendo peor que vectores estáticos promediados en similitud de frases — es lo que reproduce la tabla anterior. ↩
-
Karpukhin, V. et al. Dense Passage Retrieval for Open-Domain Question Answering. arXiv:2004.04906 (2020). Entrenamiento contrastivo de un retriever de dos encoders; el ancestro directo del stack de retrieval del capítulo 19. ↩