Capítulo 17 de 19 11 secciones 18 min

Generar dígitos nuevos en vez de reconocerlos

Un modelo de difusión escrito en numpy, en 35 segundos y sin ninguna librería de deep learning. Le pides un 7 y te dibuja un 7.

Todo lo que hicimos hasta aquí fue mirar algo y ponerle nombre. Este capítulo hace lo contrario y es el que yo tenía más ganas de escribir: le pides un 7 y te dibuja un 7 que nunca existió. Son cuarenta líneas de numpy y 35 segundos, y la idea de fondo es la misma que la de los modelos que generan imágenes hoy 🎨

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.

Un dígito de verdad al que se le echa ruido 60 veces hasta quedar en ruido puro, y desde ese ruido puro la red quita el ruido 60 veces y aparece un dígito que nunca existió.
La flecha de bajada es fácil y no hace falta aprenderla: echar ruido lo sabe hacer cualquiera. Lo único que la red aprende es a deshacer un paso, y generar es hacerlo sesenta veces seguidas.

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úmeros1024 por 1024 a color, o sea tres millones
Una red de dos capas de 256Una U-Net con atención, la del capítulo 13, de miles de millones de pesos
Le pides un dígito del 0 al 9Le pides una frase, y esa frase entra como los embeddings del capítulo 11
Se limpia el ruido sobre los píxelesSe limpia sobre una versión comprimida de la imagen, que es lo que hace que quepa en una tarjeta
60 pasos y 35 segundosDe 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 pidenQué diría yo
Fotos de producto para el catálogo sin sesión de fotosAquí sí, y con un modelo ya hecho. Nunca entrenando el tuyo
Inventar datos de ventas para probar el sistema sin usar los de clientesSe 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 facturasNada de esto. Eso es imputación y está en el libro de machine learning
Escribir las descripciones de los productos del stockUn 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! 🌸

¿Tienes alguna duda o consulta?