Enmascarar o no enmascarar: el efecto de los tokens de aviso en el ajuste de instrucciones | de David Vaughn | septiembre de 2024

Estos gráficos sugieren que cuando un conjunto de datos rg distribución cubre múltiples órdenes de magnitud o tiene una representación no despreciable tanto en Rg>1 y Rg<1 regiones (como en el caso de OpenOrca y otros conjuntos de datos con R̅g>1) la distribución puede volverse muy sesgada. Como resultado, la media aritmética puede verse influenciada desproporcionadamente por valores más grandes, lo que potencialmente tergiversa la tendencia central de la distribución. En tales casos, calcular la media en el espacio logarítmico (luego, opcionalmente, transformarla nuevamente a la escala original) podría proporcionar una estadística resumida más significativa. En otras palabras, podría tener sentido utilizar el media geométrica:

El conjunto de datos de comprensión lectora de RACE

Basado en lo anterior R̅g mesa, decidí el CARRERA RmiAtimbre docomprensión Conjunto de datos de miexámenes (R̅g=0,01) sería un buen candidato para la investigación. El control de calidad de opción múltiple parecía un banco de pruebas ideal para explorar los efectos del enmascaramiento de mensajes, ya que el mensaje es naturalmente muy largo en relación con su finalización. Independientemente de la duración del mensaje, la finalización es siempre 1 carácter de largo, es decir A, B, do o D (si ignora tokens especiales, delimitadores, etc.). Mi corazonada fue que si Si hay algún efecto al modular los pesos de los tokens de aviso, ciertamente se notarían aquí.

Como se indica en el tarjeta de conjunto de datos:

RACE es un conjunto de datos de comprensión lectora a gran escala con más de 28.000 pasajes y casi 100.000 preguntas. El conjunto de datos se recopila de los exámenes de inglés en China, que están diseñados para estudiantes de secundaria y preparatoria. El conjunto de datos puede servir como conjunto de entrenamiento y prueba para la comprensión de la máquina.

El esquema de control de calidad es simple: el mensaje presenta una preguntaposiblemente algún contexto (el artículo campo), y luego enumera cuatro opciones. La finalización (respuesta) es siempre uno de: A, B, C, D. Este visor de conjuntos de datos alojado en HuggingFace permite navegar por el conjunto completo, pero aquí hay un pequeño ejemplo:

Ejemplo de RACE (captura de pantalla de https://huggingface.co/datasets/ehovy/race/viewer/all/train)

Antes de pasar a la implementación completa de pérdida-de-peso-prontay probarlo con los datos de RACE, necesitamos una comprensión básica de la pérdida y de dónde proviene. En pocas palabras, la pérdida es una medida de qué tan bien nuestro modelo (LLM) “se ajusta” (explica, predice) nuestros datos. Durante el ajuste fino (y también el entrenamiento previo), “acercamos” el modelo a los datos ajustando los pesos de la red de tal manera que disminuya la pérdida. El regla de la cadena (de cálculo) nos proporciona un algoritmo preciso para calcular estos ajustes, dada la función de pérdida y la arquitectura de la red.

La función de pérdida más común en el ajuste fino de LLM se llama Pérdida de entropía cruzada (CEL). Por esta razón, la mayoría de las discusiones sobre CEL se enmarcan en torno a la definición de entropía cruzadaque proviene de la teoría de la información. Si bien es cierto que la “entropía cruzada” está ahí en el nombre, se puede lograr una comprensión más intuitiva al abordar CEL a través de la lente de estimación de máxima verosimilitud (MLE). Intentaré explicarlo desde ambos ángulos.

Ya hemos establecido que los LLM están programados para próxima predicción simbólica. Lo que esto significa es que el LLM es básicamente una función matemática que toma como entrada una secuencia de tokens y genera una distribución de probabilidad condicional para el siguiente token sobre todo el vocabulario simbólico V. En otras palabras, genera un vector de valores de probabilidad de dimensión |V| que suma 1. (en notación establecida |S| denota el número de elementos, o cardinalidadde un conjunto S)

Tomemos un pequeño ejemplo de juguete para ilustrar cómo funciona. Imagine que nuestros datos de entrenamiento contienen la secuencia de 4 tokens: The bird flew away. Dadas las primeras 3 fichas (The bird flew), un LLM podría generar el siguiente vector de probabilidades para cada 4ᵗʰ token posible; en aras de la simplicidad, imaginaremos que los 5 tokens candidatos enumerados (en magenta) son las únicas posibilidades (es decir, |V|=5). la funcion pag() representa las probabilidades condicionales generadas por el LLM (observe que suman 1):

(imagen del autor)

Al entrenar (o ajustar) un LLM en una secuencia de tokens, recorremos la secuencia token por token y comparamos los distribución del siguiente token generado por el LLM a la siguiente token real en la secuencia, y a partir de ahí calculamos el CEL para ese token.

Observe aquí que la ficha 4ᵗʰ real en la secuencia (away) hace no tiene la probabilidad más alta en la tabla. Durante el entrenamiento, nos gustaría ajustar ligeramente los pesos para aumentar la probabilidad de awaymientras que los demás disminuyen. El llave es tener la función de pérdida correcta… nos permite calcular exactamente cuánto modificar cada peso, para cada token.

Una vez que se calcula la pérdida para cada token, la pérdida final se calcula como el pérdida promedio por token sobre todos los tokens. Pero primero debemos establecer la fórmula para esta pérdida por token.

Interpretación de la teoría de la información.

Continuando con el problema del juguete, para calcular CEL para la posición de la ficha 4ᵗʰ, comparamos el actual 4ᵗʰ token a la distribución generada pag() sobre los 5 posible 4ᵗʰ fichas. De hecho, tratamos el token 4ᵗʰ real como una distribución. q() por derecho propio (aunque degenerado) que tiene un valor de 1 para el token que aparece en los datos –away– y un valor de 0 para todos los demás posibles tokens 4ᵗʰ (esto a veces se llama codificación one-hot).

(imagen del autor)

La razón por la que distorsionamos los datos de entrenamiento en este extraño uno-caliente representación de probabilidad codificada q() es para que podamos aplicar la fórmula para doentropía-de-rossque es una medida de la divergencia entre dos distribuciones de probabilidad discretas (por cierto, no simétricas respecto a q,p):

dónde incógnita índices sobre todos los estados posibles (es decir, 5 tokens). Esto resulta en:

Entonces, básicamente, CEL solo está usando el q vector para seleccionar entre pag vector el valor único correspondiente al token que de hecho aparece en los datos –away– (es decir, multiplicarlo por 1) y descartar todos los demás valores (es decir, multiplicar por 0). Por lo tanto, estamos indexando todos los estados posibles (tokens) solo para seleccionar uno e ignorar el resto.

Interpretación MLE

Al ajustar un LLM, buscamos los pesos θ del LLM que maximicen la probabilidad de los datos de entrenamiento dados esos pesos, a menudo llamados probabilidad de los pesos ℒ(θ) = ℙ(D|θ). Y entonces requerimos una expresión para esta cantidad. Afortunadamente, existe una manera fácil de calcular esto a partir de las probabilidades del siguiente token, que ya nos brinda el LLM.

Comenzando con el otro regla de la cadena (de probabilidad)descomponemos la probabilidad conjunta de una secuencia de tokens S en un producto de probabilidades condicionales:

Regla de la cadena (probabilidad)

Esta descomposición establece la conexión entre la predicción del siguiente token y la probabilidad conjunta de la secuencia completa del token: la probabilidad conjunta es solo el producto de todos los condicionales.

Usando i para indexar los tokens de una secuencia de tokens S = (t₁,t₂,t₃,…, tᵢ ,…)usaremos la siguiente abreviatura para denotar la salida de probabilidad condicional de un LLM para el iᵗʰ token en una secuencia, dados los pesos LLM θ y el anterior yo-1 fichas:

Cabe recalcar que pᵢ es no un vector aquí (es decir, una distribución sobre todos los posibles tokens siguientes) pero representa solo la probabilidad calculada para el evento real iᵗʰ token, es decir, la fila resaltada en amarillo en el ejemplo anterior.

Si tomamos el logaritmo de la probabilidad conjunta de una secuencia, un producto se convierte en una suma (dado que log es monótono, esto no afecta la optimización):

Ahora podemos conectar la expresión final de suma de registros (aquí mismo☝)️ a la fórmula para Pérdida promedio de entropía cruzada l sobre una secuencia de tokens:

que es la función objetivo del modelo de lenguaje causal. A menudo el “Promedio” se elimina del nombre y simplemente se llama “Pérdida de entropía cruzada”, pero es bueno recordar que CEL se calcula técnicamente a nivel de token y luego se promedia entre los tokens. De esta expresión final debería quedar claro que minimizando el CEL es equivalente a maximizar la probabilidad de la secuencia del tokenque es lo que busca MLE.

Una ventaja que resulta de la forma de esta expresión es que es muy fácil de modificar si queremos calcular la pérdida sobre cualquier subconjunto de las fichas. Recuerde que a veces podemos estar interesados ​​en encontrar los pesos θ del LLM que maximicen la probabilidad de finalización dada la indicación:

Podríamos ajustar fácilmente la pérdida para este escenario simplemente promediando solo los tokens de finalización. Si usamos “𝕀c” a denotar el conjunto de todos los índices de tokens de finalización, entonces podemos expresar pérdida de finalización como:

Dado que la pérdida de cada token ya está condicionada a todos los tokens anteriores en la secuencia, esto significa que la solicitud se contabiliza automáticamente en el condicional, incluso si solo promediamos los tokens de finalización excesiva.

Ahora que hemos establecido CEL como un promedio de pérdidas por token durante una secuencia de token, podemos definir el promedio ponderado versión de CEL:

Dependiendo de cómo establezcamos los pesos wᵢpodemos usar esta fórmula para definir pérdidas múltiples. Por ejemplo, si configuramos todos los pesos wᵢ=1 luego recuperamos el CEL de secuencia completa estándar de antes. Sin embargo, si establecemos wᵢ=1 solo para fichas de finalización, y wᵢ = 0 para tokens de aviso, entonces obtenemos pérdida de finalización. Y de la misma manera, pérdida inmediata se define estableciendo wᵢ=1 solo sobre tokens de aviso, y wᵢ = 0 de lo contrario.

Dado que rara vez (o nunca) queremos reducir el peso de los tokens de finalización, fijamos los pesos de los tokens de finalización en wᵢ=1pero para los tokens de aviso podemos definir un valor continuo en el [0:1] intervalo llamado prompt_loss_weight. De esta manera podemos ajustar cuánto pesar las fichas de aviso durante el entrenamiento, desde wᵢ = 0 (pérdida de finalización) hasta el final wᵢ=1 (pérdida de secuencia completa estándar). O incluso podríamos usar wᵢ=0.1 para dar a las fichas de aviso un peso pequeño pero distinto de cero.

Implementación de pérdidas

Echemos un vistazo más allá de cómo se calcula normalmente la pérdida en el AbrazosCara transformadores paquete. Ya que estaremos afinando el Llama-2–7b-chat-hf modelo en nuestros experimentos, veremos LlamaForCausalLMespecíficamente en el pase hacia adelantedonde la pérdida se calcula durante el entrenamiento.

Recuerde que la pérdida es una forma de comparar cada actual token para el LLM predicción para ese token (dados los tokens reales anteriores), por lo que la función de pérdida necesita acceso a estas dos estructuras de datos. En este caso, la pérdida se alimenta de dos tensores: logitsy labels. El labels tensor contiene los tokens reales (identificadores de tokens para ser exactos). Ellogits El tensor contiene las probabilidades del siguiente token predichas, antes de softmax normalización (que los obliga a sumar 1; resulta que es más eficiente dejar estos valores en su forma cruda y prenormalizada).

El logits tensor es 3D, con forma [B,N,|V|]dónde B es el tamaño del lote, N es la longitud de la secuencia (en tokens), y |V| es el tamaño del vocabulario simbólico. El 2D labels tensor solo contiene la secuencia del token en sí, por lo que tiene forma [B,N]. Aquí está la sección clave del código donde normalmente se calcula CEL:

# Shift-by-1 so that tokens < n predict n
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()

# Flatten the tensors
shift_logits = shift_logits.view(-1, self.config.vocab_size)
shift_labels = shift_labels.view(-1)

# Enable model parallelism
shift_labels = shift_labels.to(shift_logits.device)

# Compute loss
loss_fct = CrossEntropyLoss()
loss = loss_fct(shift_logits, shift_labels)

Para cada posición i a lo largo de la segunda dimensión de logitseste tensor contiene probabilidades para predecir la próximo ficha (ficha yo+1) dadas todas las fichas anteriores a través de el iᵗʰ token. Estas probabilidades deben compararse con las reales. yo+1ˢᵗ ficha en labels. Esta es la razón por la que cambio por 1 sucede en las primeras líneas: para alinear estos dos valores para cada token.