Quince capítulos mirando cosas y poniéndoles nombre. Este va al revés 🎨
Aquí no hay nada que clasificar: hay que inventar un dígito que no existe, y que igual se vea como un dígito. Y lo vamos a hacer con numpy, como todo en este libro, en cuarenta líneas y en medio minuto.
Antes de empezar, quédate con la pregunta, que es la que a mí me costó entender: ¿cómo se aprende a dibujar algo que nunca has visto? La respuesta es preciosa y es contraintuitiva 🐣
Los dígitos, y cómo se ven escritos con letras
Como aquí no puedo enseñarte imágenes, las voy a escribir con caracteres: cada píxel se convierte en un símbolo según cuánta tinta tiene. Se ve raro los primeros diez segundos y después se lee solo:
import numpy as np from sklearn.datasets import load_digits d = load_digits() X = (d.data / 16.0) * 2 - 1 # de 0..16 a -1..1 y = d.target ESCALA = ' .:-=+*#%@' def dibuja(v): # Una imagen de 8 por 8 escrita con caracteres, que es lo que cabe en un libro. im = ((v.reshape(8, 8) + 1) / 2).clip(0, 1) return [''.join(ESCALA[int(c * 9)] for c in fila) for fila in im] def lado_a_lado(imagenes, titulos): filas = [dibuja(v) for v in imagenes] print(' '.join(t.center(8) for t in titulos)) for k in range(8): print(' '.join(f[k] for f in filas)) lado_a_lado([X[0], X[1], X[2]], ['un 0', 'un 1', 'un 2'])
un 0 un 1 un 2 :#+ *#: :%* #%+%: *@+ .@%# .%. *= .%@- =#=@ :* == -%@@. -%* := += @@. =#% :* *- @@- +@@: .#:+* @@- .#@@*: -#+ *@+ .*@+
Ahí están, un cero, un uno y un dos 🙂
Fíjate también en la primera línea del código: los píxeles pasan de ir de 0 a 16 a ir de -1 a 1. Eso es la normalización de siempre, y aquí importa más que nunca porque el ruido que vamos a usar está centrado en cero.
La idea, que es de las más bonitas que hay
Aprender a dibujar de cero es dificilísimo. Pero hay algo que sí es fácil: destruir un dibujo. Le echas un poquito de ruido, y otro poquito, y otro, y a los sesenta pasos ya no queda nada.
La jugada es esta: si entrenamos una red para deshacer un paso de esa destrucción, después podemos arrancar de ruido puro y deshacer sesenta pasos seguidos. Y lo que salga al final va a ser un dígito que nadie dibujó 🌀
Eso es un modelo de difusión, y es la familia a la que pertenecen los modelos que generan imágenes que ves por ahí. El nuestro es chiquito y el de ellos es enorme, pero la cuenta es la misma.
Primero el plan de destrucción. En cada paso se conserva un poquito menos del dibujo original:
T = 60 betas = np.linspace(1e-4, 0.10, T) alfas = 1 - betas abarra = np.cumprod(alfas) print('cuanto queda del dibujo en el paso 0 :', round(float(abarra[0]), 4)) print('cuanto queda en el paso 30 :', round(float(abarra[30]), 4)) print('cuanto queda en el ultimo paso :', round(float(abarra[-1]), 4))
cuanto queda del dibujo en el paso 0 : 0.9999 cuanto queda en el paso 30 : 0.4473 cuanto queda en el ultimo paso : 0.0446
abarra es el producto acumulado de todo lo que se conservó hasta
ese paso, y es lo que permite el atajo que hace esto viable: para saber cómo se
ve una imagen en el paso 35 no hay que dar 35 pasos, se calcula
de una vez.
rng = np.random.default_rng(0) def ensucia(x0, t): ruido = rng.normal(size=x0.shape) return np.sqrt(abarra[t]) * x0 + np.sqrt(1 - abarra[t]) * ruido lado_a_lado([X[0], ensucia(X[0], 15), ensucia(X[0], 35), ensucia(X[0], 59)], ['paso 0', 'paso 15', 'paso 35', 'paso 59'])
paso 0 paso 15 paso 35 paso 59 :#+ =#-.:: - @@# := . **@-% #%+%: +% # + @.:@-@ @% @*#++ .%. *= .%- %-. : + #** =: - @-- :* == ::+ =: @=# .++ .= == := += ==. =*- :. -%- +:::@== :* *- +%..+*= + := #:: @%@-#*#: .#:+* -=% +# . .-@ @%=* % -- *: -#+ ...+= - === = : .+:@*.
Del cero del principio al puré del final 🫠
En el paso 15 todavía se adivina, en el 35 ya casi no, y en el 59 no queda nada. Y esa última columna es importante: si el final es ruido puro, entonces para generar podemos empezar por ahí, porque ruido puro sabemos fabricar.
Lo que la red aprende no es a dibujar
Aquí está el truco, y es donde a mí se me hizo el clic.
La red no aprende a dibujar dígitos. Aprende algo mucho más modesto: mirar una imagen sucia y adivinar qué ruido se le echó. Eso es un problema de regresión normalito, de los del capítulo 5, con su error cuadrático y todo.
Y si sabes qué ruido hay, sabes quitarlo. Ahí está todo.
A la red le vamos a dar tres cosas: la imagen sucia, en qué paso vamos, y qué dígito queremos. Lo del paso importa porque quitar ruido en el paso 5 y en el 55 son trabajos distintos, y se le pasa con senos y cosenos en vez de con un número suelto, que es como se hace en los modelos de verdad:
FREQ = np.array([1., 2., 4., 8.]) P, H = 64, 256 ENTRADA = P + 8 + 10 # la imagen, el reloj y el digito que se pide def entrada(xt, t, cual): ang = (t[:, None] / T) * FREQ * np.pi reloj = np.concatenate([np.sin(ang), np.cos(ang)], axis=1) return np.concatenate([xt, reloj, np.eye(10)[cual]], axis=1) def adivina_el_ruido(p, xt, t, cual): z = entrada(xt, t, cual) h1 = np.maximum(0, z @ p['W1'] + p['b1']) h2 = np.maximum(0, h1 @ p['W2'] + p['b2']) return h2 @ p['W3'] + p['b3'], (z, h1, h2) def entrena(vueltas=6000, lote=256, paso=0.002, semilla=0): r = np.random.default_rng(semilla) p = dict(W1=r.normal(0, np.sqrt(2 / ENTRADA), (ENTRADA, H)), b1=np.zeros(H), W2=r.normal(0, np.sqrt(2 / H), (H, H)), b2=np.zeros(H), W3=r.normal(0, np.sqrt(2 / H), (H, P)), b3=np.zeros(P)) m = {k: np.zeros_like(v) for k, v in p.items()} v = {k: np.zeros_like(w) for k, w in p.items()} for it in range(1, vueltas + 1): i = r.integers(0, len(X), lote) t = r.integers(0, T, lote) ruido = r.normal(size=(lote, P)) a = abarra[t][:, None] xt = np.sqrt(a) * X[i] + np.sqrt(1 - a) * ruido # se ensucia pred, (z, h1, h2) = adivina_el_ruido(p, xt, t, y[i]) # se adivina d3 = 2 * (pred - ruido) / lote # y se corrige g = {'W3': h2.T @ d3, 'b3': d3.sum(0)} d2 = (d3 @ p['W3'].T) * (h2 > 0) g['W2'] = h1.T @ d2 g['b2'] = d2.sum(0) d1 = (d2 @ p['W2'].T) * (h1 > 0) g['W1'] = z.T @ d1 g['b1'] = d1.sum(0) for k in p: # Adam, el optimizador de siempre m[k] = 0.9 * m[k] + 0.1 * g[k] v[k] = 0.999 * v[k] + 0.001 * g[k] ** 2 p[k] -= paso * (m[k] / (1 - 0.9 ** it)) / (np.sqrt(v[k] / (1 - 0.999 ** it)) + 1e-8) if it % 2000 == 0: print(f'vuelta {it:5} error al adivinar el ruido {((pred - ruido) ** 2).mean():.4f}') return p modelo = entrena()
vuelta 2000 error al adivinar el ruido 0.2213 vuelta 4000 error al adivinar el ruido 0.2000 vuelta 6000 error al adivinar el ruido 0.2094
Fíjate en que ese error no baja a cero, ni de lejos, y eso está bien 🙂
Adivinar el ruido exacto es imposible: hay infinitos ruidos que llevan a la misma imagen sucia. La red aprende el promedio de todos ellos, y ese promedio es justo lo que hace falta. Un error de 0,20 sobre un ruido de varianza 1 significa que acierta el ochenta por ciento de lo que hay que quitar.
La retropropagación de en medio es exactamente la del capítulo 6, escrita a mano igual que allí. Y el optimizador es Adam, que es lo único nuevo: en vez de restar el gradiente tal cual, guarda un promedio de por dónde iba y otro de cuánto se movía, y ajusta el paso de cada peso por separado 🧭
Y ahora se camina hacia atrás
Se arranca de ruido puro y se deshacen los sesenta pasos, uno por uno, pidiéndole el dígito que queremos:
def genera(p, pedido, semilla=7): r = np.random.default_rng(semilla) x = r.normal(size=(len(pedido), P)) # se arranca de ruido puro for t in range(T - 1, -1, -1): ruido, _ = adivina_el_ruido(p, x, np.full(len(pedido), t), pedido) x0 = ((x - np.sqrt(1 - abarra[t]) * ruido) / np.sqrt(abarra[t])).clip(-1, 1) antes = abarra[t - 1] if t > 0 else 1.0 media = (np.sqrt(antes) * betas[t] / (1 - abarra[t])) * x0 \ + (np.sqrt(alfas[t]) * (1 - antes) / (1 - abarra[t])) * x if t > 0: x = media + np.sqrt(betas[t] * (1 - antes) / (1 - abarra[t])) * r.normal(size=x.shape) else: x = media return x pedido = np.repeat(np.arange(10), 30) # 30 de cada digito inventados = genera(modelo, pedido) lado_a_lado([inventados[i * 30] for i in range(5)], ['0', '1', '2', '3', '4']) print() lado_a_lado([inventados[i * 30] for i in range(5, 10)], ['5', '6', '7', '8', '9'])
0 1 2 3 4 =#- ** .##: ## +%: .*#**. -##. %%## %+% =#=.. %= #: . #* *% #% :-:@. -:.. .%. +: : +*+ .: %= ..%%: :* -*: :@-.== =+# . +# .=%: =%-*%: :% .%- -#+. .%+ - +#. :#+%# %=+% . .**+.. .*@*=: *:#%: -+@. =*#: *%#=: *%- #+++. ++ 5 6 7 8 9 :-+*: *. +@#.. .-+#= ==+- #+-+. *# .. +#%%: =%*#- =*-=- + . .* :.%. #+=* # =%= +#*#+. :=... ::*% :+*:. %##%= .*..+: -%=:*. +#@*. *#% =::* : .=: .%-.:+ -*#: :@:+- :+ -::+ *=-++ *- = #+ :=+ :*= +%* -%: .=#:. +#.
Eso salió de ruido 😍
Mira el 0, el 1, el 3 y el 7. Ninguno de esos dibujos existía: la red arrancó de sesenta números al azar y los fue limpiando hasta que quedó eso.
Y mira también los que no salieron tan bien, el 2 y el 5, que están manchados. Con una red de dos capas y ocho por ocho píxeles esto es lo que hay, y prefiero enseñártelos que elegirte solo los bonitos 🫶
Pero mi ojo no es una medición
Que a mí me parezca un siete no vale nada. Vamos a ponerle un juez, que es una regresión logística entrenada con dígitos de verdad:
from sklearn.linear_model import LogisticRegression from sklearn.model_selection import train_test_split X_tr, X_te, y_tr, y_te = train_test_split(X, y, test_size=0.3, random_state=0, stratify=y) juez = LogisticRegression(max_iter=2000).fit(X_tr, y_tr) print('el juez acierta en digitos de verdad :', round(juez.score(X_te, y_te), 4)) print('y le pone al inventado el digito que pedi:', round(float((juez.predict(inventados) == pedido).mean()), 4)) print('confianza media en los de verdad :', round(float(juez.predict_proba(X_te).max(1).mean()), 4)) print('confianza media en los inventados :', round(float(juez.predict_proba(inventados).max(1).mean()), 4))
el juez acierta en digitos de verdad : 0.9741 y le pone al inventado el digito que pedi: 0.9733 confianza media en los de verdad : 0.941 confianza media en los inventados : 0.9163
0,9733 🎯
De 300 dígitos inventados, el juez le pone a 292 el número que yo le pedí. Y para comparar contra algo: ese mismo juez acierta 0,9741 en dígitos escritos por personas de verdad. O sea que los inventados le resultan casi tan legibles como los auténticos.
La confianza también dice algo, y dice la verdad: 0,9163 contra 0,9410. Los inventados son un pelín más dudosos. Se nota, y está bien que se note.
¿Y no estará copiando?
Esta es la pregunta que hay que hacerle siempre a un modelo generativo, y casi nadie la hace. Un modelo que se aprendiera de memoria las 1.797 imágenes y te devolviera copias sacaría 0,9733 igual, y no habría aprendido nada 🚩
Se comprueba midiendo distancias:
from sklearn.neighbors import NearestNeighbors vecino = NearestNeighbors(n_neighbors=2).fit(X) d_inv = vecino.kneighbors(inventados, n_neighbors=1)[0][:, 0] d_real = vecino.kneighbors(X, n_neighbors=2)[0][:, 1] print('distancia del inventado al digito real mas parecido:', round(float(d_inv.mean()), 4)) print('distancia de un digito real al real mas parecido :', round(float(d_real.mean()), 4))
distancia del inventado al digito real mas parecido: 2.4941 distancia de un digito real al real mas parecido : 2.0549
No copia 😌
Un dígito inventado está a 2,4941 del real más parecido, y dos dígitos reales cualesquiera están a 2,0549 entre sí. O sea que los inventados están más lejos del conjunto de entrenamiento de lo que están sus propias imágenes unas de otras. Si estuviera copiando, ese primer número sería casi cero.
Esta comprobación te la dejo subrayada porque vale para cualquier modelo generativo que te encuentres, no solo para este: antes de creerte que inventa, mide a qué distancia está de lo que vio 📏
Un modelo generativo no aprende a dibujar. Aprende a quitar ruido, y dibujar es hacerlo sesenta veces seguidas.
El error que sale el primer día
Le pides un dígito que no existe:
np.eye(10)[12]
IndexError: index 12 is out of bounds for axis 0 with size 10
El modelo sabe dibujar diez cosas, las diez que le enseñaste, y el 12 no es una de ellas. Parece una tontería y no lo es: es la limitación de fondo de todo esto. Un modelo generativo solo puede combinar lo que vio, y cuando le pides algo fuera de esas diez categorías no te avisa con un error tan amable, te devuelve algo raro y con cara de seguridad 🚩
Cómo se pasa de esto a lo que ves por ahí
La cuenta es la misma. Lo que cambia es el tamaño y dos ideas encima:
| Aquí | En un modelo de verdad |
|---|---|
| Imágenes de 8 por 8, o sea 64 números | 1024 por 1024 a color, o sea tres millones |
| Una red de dos capas de 256 | Una U-Net con atención, la del capítulo 13, de miles de millones de pesos |
| Le pides un dígito del 0 al 9 | Le pides una frase, y esa frase entra como los embeddings del capítulo 11 |
| Se limpia el ruido sobre los píxeles | Se limpia sobre una versión comprimida de la imagen, que es lo que hace que quepa en una tarjeta |
| 60 pasos y 35 segundos | De 20 a 50 pasos y meses de entrenamiento |
Y si quieres el original, el paper es Denoising Diffusion Probabilistic Models (2020). Lo que acabas de escribir es su algoritmo 1 y su algoritmo 2, sin recortes 📄
Y hay otra familia que conviene que conozcas de nombre, porque durante años fue la que mandaba: las redes generativas antagónicas (2014), donde una red dibuja y otra intenta pillarla, y las dos mejoran peleando. Son preciosas y son famosas por lo difíciles que son de entrenar.
Y en un negocio de verdad, ¿esto para qué?
Te lo aterrizo, que es la parte que a mí más me preguntan 💼
| Lo que te piden | Qué diría yo |
|---|---|
| Fotos de producto para el catálogo sin sesión de fotos | Aquí sí, y con un modelo ya hecho. Nunca entrenando el tuyo |
| Inventar datos de ventas para probar el sistema sin usar los de clientes | Se hace y se llama dato sintético. Ojo con creerte que el modelo reproduce las relaciones reales, que casi nunca las reproduce |
| Rellenar los montos que faltan en la tabla de facturas | Nada de esto. Eso es imputación y está en el libro de machine learning |
| Escribir las descripciones de los productos del stock | Un modelo de lenguaje, que es la misma familia de ideas pero para texto y está en el capítulo 14 |
La segunda fila es la que más veces me toca frenar. Generar datos falsos que se ven razonables es fácil; generar datos falsos donde la relación entre ciudad, canal y monto sea la de verdad es otra cosa, y hay que comprobarlo antes de tomar una decisión con ellos 🔍
La trampa
Un equipo entrena un modelo generativo para inventar fotos de productos y lo evalúa con un clasificador entrenado sobre las mismas fotos. Sale 0,98 y lo presentan como que el modelo aprendió a dibujar productos nuevos.
modelo = entrena_generativo(fotos_del_catalogo) inventadas = modelo.genera(500) juez = entrena_clasificador(fotos_del_catalogo, categorias) print(accuracy_score(categorias_pedidas, juez.predict(inventadas))) # 0.98
Qué está mal
El 0,98 no distingue entre un modelo que inventa y un modelo que copia. Si el generativo se aprendió de memoria las fotos del catálogo y devuelve copias ligeramente movidas, el juez las va a reconocer mejor todavía, porque las vio en su propio entrenamiento. La medición que falta es la del capítulo: cuánto se parece cada imagen inventada a la imagen real más cercana, comparado con cuánto se parecen las reales entre sí. Aquí ese número es 2,4941 contra 2,0549 y por eso se puede afirmar que no copia. Sin esa comprobación, un acierto altísimo del juez es exactamente lo que uno esperaría de un modelo que no aprendió nada, y también de uno que aprendió todo.
Ejercicios
Seis. El 2 es el que más te va a enseñar de cómo funciona esto por dentro 💛
1. Pídele treinta veces el mismo dígito
Genera treinta sietes y míralos todos.
sietes = genera(modelo, np.full(6, 7), semilla=3) lado_a_lado(list(sietes), ['a', 'b', 'c', 'd', 'e', 'f'])
La pregunta que contesta: ¿son sietes distintos o es el mismo seis veces? Si salieran todos iguales tendrías un colapso, que es el fallo clásico de los modelos generativos y tiene nombre propio.
2. Mira la película de la limpieza
Guarda x en los pasos 59, 45, 30, 15 y 0, y
dibújalos en fila.
Es el ejercicio que hace que entiendas esto de verdad. Vas a ver cómo del ruido va apareciendo primero una mancha con la forma general, después el trazo, y al final los detalles. Y de paso vas a entender por qué los modelos de verdad generan primero la composición y después la textura 🎬
3. Menos pasos
Baja T a 20 y vuelve a entrenar y a generar.
Menos pasos es más rápido y cada paso tiene que hacer más trabajo. Mide las dos cosas: cuánto tarda y qué dice el juez. Esa es exactamente la tensión que tienen los modelos comerciales, y el motivo de que se investigue tanto en generar con veinte pasos en vez de con mil.
4. Quítale la etiqueta
Entrena sin pasarle qué dígito es y mira qué sale.
Se hace poniendo np.zeros((len(cual), 10)) en vez del
np.eye(10)[cual]. El modelo va a seguir generando dígitos, pero ya
no vas a poder pedirle cuál. Sirve para ver qué aporta exactamente esa parte de
la entrada, que es la diferencia entre un modelo que dibuja y uno al que le
puedes pedir cosas.
5. Quítale el reloj
Pásale t / T como un solo número en vez de los
ocho senos y cosenos.
Esto lo probé yo mientras escribía el capítulo y por eso te lo dejo: con un solo número el modelo entrena igual de rápido y genera bastante peor. La red necesita distinguir bien en qué momento está, y ocho columnas que oscilan a frecuencias distintas se lo ponen mucho más fácil que una sola que sube despacio 🕐
6. Sube el ruido máximo
Cambia 0.10 por 0.25 en las betas
y mira qué pasa.
Otro que probé y salió mal, que para eso están: con betas grandes la destrucción es demasiado brusca, cada paso hacia atrás tiene que adivinar demasiado y lo que sale son manchas. El plan de destrucción no es un detalle de implementación, es la mitad del modelo.
Comprueba que lo tienes
Tu modelo generativo saca 0,97 en el juez. ¿Qué compruebas antes de decir que funciona?
- A qué distancia está cada imagen inventada de la imagen real más parecida
- Que el juez esté bien entrenado
- Que las imágenes se vean bien a simple vista
- Que el error de entrenamiento haya bajado a cero
Lo que te llevas
- 🌀 Un modelo de difusión aprende a quitar ruido, no a dibujar. Dibujar es quitarlo sesenta veces seguidas.
- 🎯 El juez le pone al inventado el dígito que se le pidió el 0,9733 de las veces, contra 0,9741 que acierta en dígitos de verdad.
- 📏 Y no copia: 2,4941 de distancia al real más parecido, contra 2,0549 que se separan dos reales entre sí.
- 🕐 Al modelo hay que decirle en qué paso va, y con senos y cosenos funciona mucho mejor que con un número suelto.
- ⚙️ Son cuarenta líneas de numpy y 35 segundos. Lo que cambia en los modelos grandes es el tamaño, no la idea.
El capítulo 18 cierra el libro montando de principio a fin la red que reconoce estos mismos dígitos, y en el glosario de IA está el vocabulario suelto 📖
Que tengas lindo día! 🌸