Leyes de escala: Kaplan y Chinchilla

Leyes de escala: Kaplan y Chinchilla

28 min de lectura

La lección sobre la perplejidad midió al mini-GPT y le encontró una grieta: sobre las páginas de Marianela que no ha leído paga unos 3.2 nats por token, bastante más que sobre las que le sirvieron para entrenar. Detrás de esa medida hay una factura que el bloque todavía no ha sumado. El checkpoint es el paso 1 750 de un entrenamiento de 32 ventanas de 64 tokens por paso: 3 584 0003\,584\,000 tokens leídos por 136 448136\,448 parámetros, cuando la parte de entrenamiento de la novela sólo tiene 125 789125\,789. El mini-GPT leyó cada uno de ellos, de media, 28 veces y media. Con ese mismo gasto se podía haber pagado otra cosa, un mini-GPT mayor que leyera menos o uno menor que leyera más, y falta saber qué habría salido mejor.

La misma pregunta, con presupuestos once órdenes de magnitud mayores, decide cómo se entrena un modelo grande. En 2020, Kaplan y sus coautores, en OpenAI, encontraron que multiplicar por diez los parámetros, los tokens leídos o el cálculo multiplica la pérdida, cada vez, por un mismo factor, y concluyeron que un presupuesto mayor debía ir sobre todo a parámetros. GPT-3 y Gopher se entrenaron en esa línea: 175 000 y 280 000 millones de parámetros, y 300 000 millones de tokens cada uno, menos de dos por parámetro. En 2022, Hoffmann y sus coautores, en DeepMind, repitieron la medida con más de 400 modelos y llegaron a otro reparto: a partes iguales, unos 20 tokens por parámetro. Para comprobarlo entrenaron Chinchilla, 70 000 millones de parámetros y 1.4 billones (millones de millones) de tokens, con el cálculo que había costado Gopher, y ganó a Gopher y a GPT-3. Esta lección deriva de dónde sale un reparto así y le hace la cuenta al mini-GPT.

Antes de las fórmulas, las formas. La primera vista del explorable pone la pérdida frente a los parámetros, los tokens o el cálculo, con los dos ejes logarítmicos y los ajustes de los dos artículos. La de Kaplan es una recta: mueve el punto y verás el mismo 0.839 por cada multiplicación por diez de los parámetros, estés donde estés. La de Chinchilla se curva hacia un suelo, y cada multiplicación compra menos que la anterior.

La segunda vista fija el presupuesto. Con el cálculo fijo, un modelo mayor lee menos tokens, y la pérdida frente al tamaño es un valle: a la izquierda, modelos pequeños que leen mucho; a la derecha, grandes que apenas leen. Su fondo es el mejor reparto, y se desplaza por la recta de abajo cuando cambias el presupuesto. Con el botón del presupuesto de Gopher, un modelo del tamaño de Chinchilla queda en el fondo, y uno del de Gopher, en la pared de la derecha.

En la primera vista, la pérdida frente a los parámetros, los tokens o el cálculo en ejes logarítmicos: Kaplan es una recta, Chinchilla se curva hacia un suelo. En la segunda, con el cálculo fijo, el valle de la pérdida frente al tamaño y la recta que sigue su fondo, con GPT-3, Gopher, Chinchilla y el mini-GPT.

Una ley de potencia es una recta en escala logarítmica

Llamemos NN a los parámetros de un modelo, todos, embeddings incluidos; DD a los tokens que lee durante el entrenamiento, contando dos veces el que lee dos; y CC al cálculo que gasta en ello, en operaciones de coma flotante (floating-point operations, FLOPs), cada multiplicación y cada suma una. Son las letras de Kaplan, y en el curso anterior eran otras cosas (documentos, corpus, caracteres) que aquí no vuelven a aparecer.

Kaplan y sus coautores entrenaron modelos de entre 768 y 1 500 millones de parámetros (sin los embeddings, que ellos dejan fuera de NN) sobre páginas web en inglés, y encontraron que, cuando una sola de las cantidades limita y las otras sobran, la pérdida sobre texto no visto es

L(N)=(NcN)αN,L(D)=(DcD)αD,\mathcal{L}(N) = \left(\frac{N_c}{N}\right)^{\alpha_N}, \qquad \mathcal{L}(D) = \left(\frac{D_c}{D}\right)^{\alpha_D},

con αN≈0.076\alpha_N \approx 0.076 y αD≈0.095\alpha_D \approx 0.095, y una tercera igual para CC, de exponente 0.0500.050. (Los artículos escriben la pérdida LL; aquí LL es el número de capas.) Una cantidad que es una constante por una potencia de otra sigue una ley de potencia, y una ley de potencia de la pérdida es una ley de escala (scaling law). La ley de Heaps del curso anterior, la de los tipos que siguen apareciendo al alargar el corpus, era de la misma familia. Tomemos logaritmos en la primera:

log⁡L(N)=αNlog⁡Nc−αNlog⁡N,\log \mathcal{L}(N) = \alpha_N \log N_c - \alpha_N \log N,

que en los ejes (log⁡N,log⁡L)(\log N, \log \mathcal{L}) es una recta de pendiente −αN-\alpha_N. Por eso estas leyes se dibujan con los dos ejes logarítmicos, y por eso se miden así: con la pérdida de unos cuantos modelos de tamaños distintos se ajusta una recta a los logaritmos, y su pendiente es el exponente. Y por eso lo que compra cada multiplicación de NN es siempre lo mismo. Multiplicarlo por kk multiplica la pérdida por k−αNk^{-\alpha_N}, en cualquier punto: por 0.949 al duplicarlo, por 0.839 al multiplicarlo por diez. Para bajar la pérdida a la mitad habría que multiplicar NN por 21/0.0762^{1/0.076}, unas 9 000 veces.

Un suelo, y lo que miden las constantes

Una recta así no puede seguir para siempre. Si NN creciera sin límite, L(N)\mathcal{L}(N) tendería a cero, y ningún modelo predice sin error un texto que tiene algo de impredecible. Hoffmann y sus coautores le ponen suelo:

L(N,D)=L∞+ANNαN+ADDαD.\mathcal{L}(N, D) = \mathcal{L}_{\infty} + \frac{A_N}{N^{\alpha_N}} + \frac{A_D}{D^{\alpha_D}}.

L∞\mathcal{L}_{\infty} es la pérdida irreducible, el suelo del explorable: lo que quedaría con NN y DD infinitos. Por encima hay dos leyes de potencia, una que se paga por tener pocos parámetros y otra por haber leído pocos tokens. (El artículo escribe EE, AA, BB, α\alpha y β\beta; aquí BB es el tamaño del batch, y β\beta la reserva el bloque siguiente.) Sus exponentes son del orden de 0.3, y no se comparan con los 0.076 de Kaplan, porque miden sólo lo que queda por encima del suelo. Y la suma ya no es una recta en ejes logarítmicos: es la curva del explorable, que se tuerce hacia L∞\mathcal{L}_{\infty}.

Queda una advertencia, la de la perplejidad: una pérdida por token mide también al tokenizador. NcN_c, DcD_c y L∞\mathcal{L}_{\infty} cambian con el corpus y el tokenizador, y Kaplan lo avisa: otro tokenizador multiplica la pérdida por un factor, y ese factor se lo quedan las constantes, no los exponentes. Por eso el mini-GPT no tiene punto en la primera vista: sus 3.2 nats son por token de su tokenizador sobre una novela en español, y la recta de Kaplan, por token del de GPT-2 sobre páginas web en inglés.

Lo que cuesta entrenar, y cómo repartirlo

El cálculo se puede contar. La lección sobre la caché contó lo que cuesta pasar una fila por la red, cfilac_{\text{fila}}: una multiplicación por cada peso de cada matriz, que en el mini-GPT son 131 072131\,072 de sus 136 448136\,448 parámetros. Con su suma son dos operaciones por peso, así que la ida cuesta unas 2N2N por token. La vuelta cuesta el doble, y backpropagation dice por qué: por cada producto de la ida hace dos, uno que lleva el error hacia atrás con la traspuesta de la matriz y otro que forma el gradiente de sus pesos, los dos del tamaño del de la ida. Tres productos por matriz y dos operaciones por multiplicación dan

C≈6ND.C \approx 6ND.

La cuenta deja fuera la atención, las 2Lt⋅dmodel2Lt \cdot d_{\text{model}} multiplicaciones por token que dependen de la posición: en las ventanas de 64 del mini-GPT suman a lo sumo un 12 % más, y en un modelo grande, con miles de dimensiones, unos pocos por ciento. Chinchilla, que la cuenta entera, no se aparta de 6ND6ND más de un 10 %. Para el mini-GPT, C=6⋅136 448⋅3 584 000≈2.9×1012C = 6 \cdot 136\,448 \cdot 3\,584\,000 \approx 2.9 \times 10^{12} FLOPs; GPT-3 gastó unos 3×10233 \times 10^{23}.

El fondo del valle

Fijemos un presupuesto, un CC decidido de antemano. Con CC dado, D=C/(6N)D = C/(6N), y la pérdida de Chinchilla pasa a depender sólo de NN: es el valle del explorable, que el artículo llama perfil IsoFLOP. Busquemos su fondo moviendo NN en escala logarítmica, que es como se mueve en el explorable. Como log⁡N+log⁡D=log⁡(C/6)\log N + \log D = \log(C/6), subir log⁡N\log N un poco baja log⁡D\log D lo mismo. Y la derivada de ANN−αNA_N N^{-\alpha_N} respecto a log⁡N\log N es −αNANN−αN-\alpha_N A_N N^{-\alpha_N}: el exponente por la potencia. Así que el término de los parámetros baja a ritmo αNAN/NαN\alpha_N A_N/N^{\alpha_N}, el de los datos sube a ritmo αDAD/DαD\alpha_D A_D/D^{\alpha_D}, y en el fondo los dos ritmos se igualan:

αNANNαN=αDADDαD.\alpha_N \frac{A_N}{N^{\alpha_N}} = \alpha_D \frac{A_D}{D^{\alpha_D}}.

Se lee sin despejar nada. En el fondo del valle, lo que baja el término de los parámetros al darle a NN un 1 % más es lo que sube el de los datos al quitarle ese 1 % a DD. Despejando con D=C/(6N)D = C/(6N) (el detalle, abajo),

N⋆∝CαD/(αN+αD),D⋆∝CαN/(αN+αD),N^{\star} \propto C^{\alpha_D/(\alpha_N + \alpha_D)}, \qquad D^{\star} \propto C^{\alpha_N/(\alpha_N + \alpha_D)},

donde el par (N⋆,D⋆)(N^{\star}, D^{\star}) es el fondo del valle, el reparto óptimo (compute-optimal) de ese presupuesto. La estrella marca el óptimo, como en las lecciones anteriores; los artículos escriben NoptN_{\text{opt}}. Los dos exponentes suman 1, como tiene que ser si N⋆D⋆=C/6N^{\star}D^{\star} = C/6. El cociente es la cuenta que importa:

D⋆N⋆∝C(αN−αD)/(αN+αD).\frac{D^{\star}}{N^{\star}} \propto C^{(\alpha_N - \alpha_D)/(\alpha_N + \alpha_D)}.

Si los dos términos bajan igual de deprisa, αN=αD\alpha_N = \alpha_D, el exponente es cero y la proporción de tokens por parámetro es la misma para cualquier presupuesto. Ésa es la afirmación de Chinchilla dicha con exponentes: N⋆N^{\star} y D⋆D^{\star} crecen los dos como C1/2C^{1/2}. En el ajuste del explorable (el del artículo, rehecho en 2024 porque el publicado se detuvo antes de tiempo), αN=0.348\alpha_N = 0.348 y αD=0.366\alpha_D = 0.366. Cuánto vale la proporción no sale de la derivación, porque depende de ANA_N y ADA_D: hay que medirla.

Ver el despeje de los dos exponentes

Sustituyamos D=C/(6N)D = C/(6N) en la condición del fondo y juntemos las potencias de NN:

αNANN−αN=αDAD(6NC)αD⟹NαN+αD=αNANαDAD(C6)αD.\alpha_N A_N N^{-\alpha_N} = \alpha_D A_D \left(\frac{6N}{C}\right)^{\alpha_D} \quad\Longrightarrow\quad N^{\alpha_N + \alpha_D} = \frac{\alpha_N A_N}{\alpha_D A_D}\left(\frac{C}{6}\right)^{\alpha_D}.

Elevando a 1/(αN+αD)1/(\alpha_N + \alpha_D) y volviendo a D=C/(6N)D = C/(6N),

N⋆=G(C6)αD/(αN+αD),D⋆=1G(C6)αN/(αN+αD),G=(αNANαDAD)1/(αN+αD),N^{\star} = G\left(\frac{C}{6}\right)^{\alpha_D/(\alpha_N + \alpha_D)}, \qquad D^{\star} = \frac{1}{G}\left(\frac{C}{6}\right)^{\alpha_N/(\alpha_N + \alpha_D)}, \qquad G = \left(\frac{\alpha_N A_N}{\alpha_D A_D}\right)^{1/(\alpha_N + \alpha_D)},

que es la ecuación 4 del artículo. El exponente de D⋆D^{\star} es 1−αD/(αN+αD)1 - \alpha_D/(\alpha_N + \alpha_D).

Kaplan midió otra cosa: N⋆∝C0.73N^{\star} \propto C^{0.73} y D⋆∝C0.27D^{\star} \propto C^{0.27}, así que por cada diez veces más de cálculo, 5.4 veces más parámetros y 1.9 veces más tokens. El artículo de Chinchilla lo achaca a que Kaplan entrenó todos sus modelos con el mismo calendario y leyó la pérdida de los entrenamientos cortos a mitad del coseno de la lección sobre el entrenamiento, antes de que la tasa bajara: salían peores de lo que eran, y con poco presupuesto los modelos grandes, que llegan a él con pocos tokens leídos, parecían peor negocio. El óptimo de los presupuestos pequeños salía menor de lo que era, y la recta de N⋆N^{\star}, más empinada. Un trabajo de 2024 señala además los embeddings: Kaplan no los contaba, y en un modelo pequeño son mucho. En el mini-GPT, 36 86436\,864 de sus 136 448136\,448 parámetros.

Las cuentas del mini-GPT

Falta el número. La celda no carga el mini-GPT: sólo necesita sus parámetros y los tokens que había leído el checkpoint. Lo que trae escrito es la tabla 3 del artículo de Chinchilla: nueve tamaños de modelo, de 400 millones a diez billones de parámetros, con los tokens que harían óptimo a cada uno según el primero de sus tres métodos, un ajuste a cientos de entrenamientos. La celda les ajusta en logaritmos las dos rectas de la derivación y las prolonga hasta el presupuesto del mini-GPT.

import numpy as np

# Hoffmann et al. (2022), tabla 3: nueve tamaños y los tokens que los hacen óptimos
N = np.array([0.4, 1, 10, 67, 175, 280, 520, 1_000, 10_000]) * 1e9
D = np.array([8.0, 20.2, 205.1, 1_500, 3_700, 5_900, 11_000, 21_200, 216_200]) * 1e9
C = 6 * N * D # el cálculo de cada uno, en FLOPs

# Las dos rectas de la derivación, ajustadas por mínimos cuadrados a los logaritmos
eN, bN = np.polyfit(np.log(C), np.log(N), 1) # log N* = eN log C + bN
eD, bD = np.polyfit(np.log(C), np.log(D), 1) # log D* = eD log C + bD
print("exponente de C: en N* %.3f, en D* %.3f, suma %.3f" % (eN, eD, eN + eD))

def reparto(c):
"""N* y D* según las dos rectas, para un presupuesto de c FLOPs."""
return np.exp(bN) * c ** eN, np.exp(bD) * c ** eD

for c in [C.min(), C.max()]:
n_opt, d_opt = reparto(c)
print("C = %.1e: N* = %.2e, D* = %.2e, D*/N* = %.1f" % (c, n_opt, d_opt, d_opt / n_opt))

# El mini-GPT: 136 448 parámetros; el checkpoint es el paso 1 750, de 32 ventanas de 64 tokens
n, d = 136_448, 1_750 * 32 * 64
distintos = 125_789 # los tokens de la parte de entrenamiento
c = 6 * n * d
n_opt, d_opt = reparto(c)
print()
print("mini-GPT: C = %.2e FLOPs, %.1f órdenes de magnitud por debajo de la tabla"
% (c, np.log10(C.min() / c)))
print(" N D D/N")
print(" el suyo %7d %9d %5.1f" % (n, d, d / n))
print(" la recta %7.0f %9.0f %5.1f" % (n_opt, d_opt, d_opt / n_opt))
print("tokens distintos: %d, leídos %.1f veces cada uno" % (distintos, d / distintos))
numpy

La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador y se reutiliza en todas las lecciones.

Los exponentes salen 0.4980.498 y 0.5020.502, un medio los dos. Que sumen 1 no comprueba nada: cada fila cumple log⁡N+log⁡D=log⁡C−log⁡6\log N + \log D = \log C - \log 6, y una recta de mínimos cuadrados respeta esa suma. La proporción sí es un resultado. A lo largo de nueve órdenes de magnitud de cálculo, las dos rectas dan entre 20.220.2 y 21.821.8 tokens por parámetro: los 20 de Chinchilla, con un 10 % de margen.

El mini-GPT queda casi siete órdenes de magnitud por debajo de la tabla, donde la recta ya no se mide sino que se prolonga (la parte discontinua del explorable, sobre la que cae su punto): la segunda cifra de lo que da no está garantizada, y el orden de magnitud sí. Para su presupuesto pide unos 160 000160\,000 parámetros y 3.13.1 millones de tokens, y el checkpoint tiene 136 448136\,448 y leyó 3.63.6 millones. En cálculo, el mini-GPT es un modelo óptimo.

Pero la recta cuenta tokens leídos, y en los experimentos de Chinchilla casi todo lo leído era texto nuevo para el modelo: su DD son, en la práctica, tokens distintos. Los 3.63.6 millones del mini-GPT son 125 789125\,789 leídos 28 veces y media. Ésa es la respuesta a la pregunta del principio: ni mayor ni menor. Con el cálculo que gastó, su tamaño es el bueno, y lo que le falta es texto, unas 24 veces más del que tiene la novela. Y releer no es leer. Una época es una pasada entera por el texto de entrenamiento, como en el descenso de gradiente del curso anterior, y un estudio de 2023 encontró que hasta unas cuatro valen casi como texto nuevo y que, a partir de ahí, cada una vale menos. Las 28 del mini-GPT están muy lejos de ese margen, y la distancia entre su pérdida de entrenamiento y la del texto reservado es lo que se ve de ello.

Comprueba tu intuición

Cinco preguntas: una ley de potencia a mano, el 6 de 6ND6ND, un reparto con exponentes desiguales, lo que dicen las cuentas del mini-GPT y un punto que no está en la gráfica.

Con la ley de Kaplan, L(N)=(Nc/N)αN\mathcal{L}(N) = (N_c/N)^{\alpha_N} con αN=0.076\alpha_N = 0.076, ¿por cuánto hay que multiplicar NN para que la pérdida baje un 10 %? Redondea a un decimal.

Se acepta un margen de ±0.1.

¿De dónde sale el 6 de C≈6NDC \approx 6ND?

Imagina una ley como la de Chinchilla en la que el término de los parámetros bajara el doble de deprisa que el de los datos: αN=2αD\alpha_N = 2\alpha_D. Con cada vez más cálculo, ¿cómo crecen N⋆N^{\star} y D⋆D^{\star}?

Marca lo que dicen las cuentas de la celda sobre el mini-GPT.

Marca todas las opciones correctas. Se corrige todo o nada: no hay puntuación parcial.

Sin contar los embeddings, el mini-GPT tiene unos 98 00098\,000 parámetros, dentro de lo que Kaplan midió, y la recta de Kaplan da para ese tamaño una pérdida de 4.84.8 nats por token. El mini-GPT paga 3.23.2 sobre su texto reservado. ¿Qué concluyes?


Las leyes de esta lección predicen un número, la pérdida sobre texto no visto, y lo predicen a lo largo de muchos órdenes de magnitud: baja sin saltos, un poco con cada multiplicación del presupuesto. Lo que no dicen es qué sabe hacer un modelo con una pérdida dada. Entre el mini-GPT y GPT-3 hay once órdenes de magnitud de cálculo, y la diferencia no se agota en que uno escriba frases de Galdós con faltas y el otro no.

El artículo de GPT-3 lleva esa diferencia en el título: Language Models are Few-Shot Learners. Con unos cuantos ejemplos de una tarea en el prompt, una para la que nadie lo entrenó (traducir, ordenar las letras de una palabra, sumar), hace la tarea con el ejemplo siguiente, sin que cambie un solo peso. Qué significa eso en la notación de este bloque, y por qué el mini-GPT no puede hacerlo, es la lección siguiente, sobre el aprendizaje en contexto.

¿Te ha sido útil?
Para profundizar5 fuentes · 5 papers

De dónde sale lo de esta lección, y dónde seguir si quieres más. Nada de aquí hace falta para continuar el curso.

  • Scaling Laws for Neural Language Models
    paperKaplan, McCandlish, Henighan, Brown, Chess, Child, Gray, Radford, Wu y Amodei, 2020arXiv:2001.08361EN

    Las leyes de potencia de esta lección, con sus constantes en la tabla 5 y el aviso de que dependen del tokenizador. Su sección 6 es el reparto que Chinchilla corrigió; la figura 1 enseña cuánto abarcan.

  • Training Compute-Optimal Large Language Models
    paperHoffmann, Borgeaud, Mensch y otros, 2022NeurIPS 2022 · arXiv:2203.15556EN

    Chinchilla. Sus tres métodos en la sección 3, el valle en la figura 3, la tabla 3 que ajusta la celda y, en el apéndice F, la cuenta exacta de FLOPs frente a 6ND.

  • Chinchilla Scaling: A replication attempt
    paperBesiroglu, Erdil, Barnett y You, 2024arXiv:2404.10102EN

    Rehacen el tercer método de Chinchilla con los datos de su figura 4 y encuentran que el ajuste publicado se paró antes de tiempo. Sus constantes corregidas son las del explorable, y dan unos 20 tokens por parámetro.

  • Reconciling Kaplan and Chinchilla Scaling Laws
    paperPearce y Song, 2024TMLR 2024 · arXiv:2406.12907EN

    Por qué Kaplan midió 0.73 y Chinchilla 0.5: Kaplan no contaba los embeddings y medía modelos pequeños, en los que pesan mucho. En el mini-GPT son más de la cuarta parte de los parámetros.

  • Scaling Data-Constrained Language Models
    paperMuennighoff, Rush, Barak, Le Scao, Piktus, Tazi, Pyysalo, Wolf y Raffel, 2023NeurIPS 2023 · arXiv:2305.16264EN

    Qué pasa cuando el texto se acaba, que es el caso del mini-GPT: hasta unas cuatro épocas, repetir vale casi como texto nuevo, y a partir de ahí cada época vale menos. Proponen una ley de escala que cuenta las repeticiones.