GRASP es un nuevo planificador basado en gradientes para dinámicas aprendidas (un "modelo mundial") que hace que la planificación a largo plazo sea práctica al (1) elevar la trayectoria a estados virtuales para que la optimización sea paralela a lo largo del tiempo, (2) agregar estocasticidad directamente a las iteraciones de estado para la exploración y (3) remodelar los gradientes para que las acciones obtengan señales limpias mientras evitamos gradientes frágiles de "entrada de estado" a través de modelos de visión de alta dimensión.
Los grandes modelos mundiales aprendidos son cada vez más capaces. Pueden predecir largas secuencias de observaciones futuras en espacios visuales de alta dimensión y generalizar tareas de maneras que eran difíciles de imaginar hace unos años. A medida que estos modelos escalan, comienzan a parecerse menos a predictores de tareas específicas y más a simuladores de propósito general.
Pero tener un modelo predictivo potente no es lo mismo que poder utilizarlo eficazmente para control/aprendizaje/planificación. En la práctica, la planificación a largo plazo con modelos del mundo moderno sigue siendo frágil: la optimización se vuelve mal condicionada, la estructura no codiciosa crea mínimos locales malos y los espacios latentes de alta dimensión introducen modos de falla sutiles.
En esta publicación de blog, describo los problemas que motivaron este proyecto y nuestro enfoque para abordarlos: por qué la planificación con modelos del mundo moderno puede ser sorprendentemente frágil, por qué los horizontes largos son la verdadera prueba de resistencia y qué cambiamos para hacer que la planificación basada en gradientes sea mucho más sólida.
Esta publicación de blog analiza el trabajo realizado con Mike Rabbat, Aditi Krishnapriyan, Yann LeCun y Amir Bar (* indica igualdad de asesoramiento), donde proponemos GRASP.
¿Qué es un modelo mundial?
Hoy en día, el término "modelo mundial" está bastante sobrecargado y, dependiendo del contexto, puede significar un modelo dinámico explícito o algún estado interno implícito y confiable en el que se basa un modelo generativo (por ejemplo, cuando un LLM genera movimientos de ajedrez, si hay alguna representación interna del tablero). A continuación damos nuestra definición de trabajo flexible.
Supongamos que realiza acciones $a_t in mathcal{A}$ y observa estados $s_t in mathcal{S}$ (imágenes, vectores latentes, propiocepción). Un modelo mundial es un modelo aprendido que, dado el estado actual y una secuencia de acciones futuras, predice lo que sucederá a continuación. Formalmente, define una distribución predictiva en una secuencia de estados observados $s_{th:t}$ y la acción actual $a_t$:
[P_theta(s_{t+1} mid s_{th:t},; a_t)]
que se aproxima al verdadero condicional $P(s_{t+1} mid s_{th:t},; a_t)$ del entorno. Para esta publicación de blog, asumiremos un modelo Markoviano $P(s_{t+1} mid s_{th:t},; a_t)$ por simplicidad (todos los resultados aquí se pueden extender al caso más general), y cuando el modelo es determinista se reduce a un mapa sobre estados:
[s_{t+1} = F_theta(s_t, a_t).]
En la práctica, el estado $s_t$ es a menudo una representación latente aprendida (por ejemplo, codificada a partir de píxeles), por lo que el modelo opera en un espacio (teóricamente) compacto y diferenciable. El punto clave es que un modelo mundial proporciona un simulador diferenciable; puede avanzar bajo secuencias de acción hipotéticas y propagarlo hacia atrás a través de las predicciones.
Planificación: elegir acciones optimizando a través del modelo.
Dado un inicio $s_0$ y una meta $g$, el planificador más simple elige una secuencia de acción $mathbf{a}=(a_0,dots,a_{T-1})$ implementando el modelo y minimizando el error terminal:
[min_{mathbf{a}} ; | s_T(mathbf{a}) – g |_2^2, quad text{donde } s_T(mathbf{a}) = mathcal{F}_{theta}^{T}(s_0,mathbf{a}).]
Aquí usamos $mathcal{F}^T$ como abreviatura para el despliegue completo a través del modelo mundial (la dependencia de los parámetros del modelo $theta$ está implícita):
[mathcal{F}_{theta}^{T}(s_0, mathbf{a}) = F_theta(F_theta(cdots F_theta(s_0, a_0), cdots, a_{T-2}), a_{T-1}).]
En horizontes cortos y sistemas de pocas dimensiones, esto puede funcionar razonablemente bien. Pero a medida que los horizontes crecen y los modelos se hacen más grandes y expresivos, sus debilidades se amplifican.
Entonces, ¿por qué esto no funciona simplemente a escala?
Por qué es difícil planificar a largo plazo (incluso cuando todo es diferenciable)
Hay dos puntos débiles separados para el modelo mundial más general, más un tercero que es específico de los modelos aprendidos y basados en el aprendizaje profundo.
1) Las implementaciones a largo plazo crean gráficos computacionales profundos y mal condicionados
Aquellos que estén familiarizados con la retropropiedad a través del tiempo (BPTT) pueden notar que nos estamos diferenciando a través de un modelo aplicado a sí mismo repetidamente, lo que conducirá al problema de los gradientes explosivos/desaparecidos. Es decir, si tomamos derivadas (tenga en cuenta que estamos diferenciando funciones con valores vectoriales, lo que da como resultado jacobianos que denotamos con $D_x (cdots)$) con respecto a acciones anteriores (por ejemplo, $a_0$):
[D_{a_0} mathcal{F}_{theta}^{T}(s_0, mathbf{a}) = Bigl(prod_{t=1}^T D_s F_theta(s_t, a_t)Bigr) D_{a_0}F_theta(s_0, a_0).]
Vemos que el condicionamiento jacobiano escala exponencialmente con el tiempo $T$:
[sigma_{text{max/min}}(D_{a_0}mathcal{F}_{theta}^{T}) sim sigma_{text{max/min}}(D_s F_theta)^{T-1},]
lo que lleva a gradientes explosivos o que desaparecen.
2) El paisaje no es codicioso y está lleno de trampas.
En horizontes cortos, la solución codiciosa, en la que avanzamos directamente hacia la meta a cada paso, suele ser suficiente. Si solo necesita planificar unos pocos pasos por delante, la trayectoria óptima generalmente no se desvía mucho de “dirigirse hacia $g$” en cada paso.
A medida que los horizontes crecen, suceden dos cosas. En primer lugar, es más probable que las tareas más largas requieran un comportamiento no codicioso: rodear una pared, reposicionarse antes de empujar, retroceder para tomar un camino mejor. Y a medida que crecen los horizontes, normalmente se necesitan más medidas no codiciosas. En segundo lugar, el espacio de optimización en sí escala con el horizonte: $mathrm{dim}(mathcal{A} times cdots times mathcal{A}) = Tmathrm{dim}(mathcal{A})$, expandiendo aún más el espacio de mínimos locales para el problema de optimización.
Una solución a largo plazo: eliminar la restricción dinámica
Supongamos que tratamos la restricción dinámica $s_{t+1} = F_{theta}(s_t, a_t)$ como una restricción suave y, en su lugar, optimizamos la siguiente función de penalización sobre ambas acciones $(a_0,ldots,a_{T-1})$ y estados $(s_0,ldots,s_T)$:
[min_{mathbf{s},mathbf{a}} mathcal{L}(mathbf{s}, mathbf{a}) = sum_{t=0}^{T-1} big|F_theta(s_t,a_t) – s_{t+1}big|_2^2, quad text{con } s_0 text{ fijo y } s_T=g.]
A esto también se le llama a veces colocación en la literatura de planificación/robótica. Tenga en cuenta que la formulación modificada comparte los mismos minimizadores globales que el objetivo de implementación original (ambos son cero exactamente cuando la trayectoria es dinámicamente factible). Pero los panoramas de optimización son muy diferentes y obtenemos dos beneficios inmediatos:
Cada evaluación del modelo mundial $F_{theta}(s_t,a_t)$ depende solo de variables locales, por lo que todos los términos $T$ se pueden calcular en paralelo a lo largo del tiempo, lo que resulta en una enorme aceleración para horizontes más largos, y ya no se propaga hacia atrás a través de una sola composición profunda de $T$ pasos para obtener una señal de aprendizaje, ya que el producto anterior de los jacobianos ahora se divide en una suma, por ejemplo: [D_{a_0} mathcal{L} = 2(F_theta(s_0, a_0) – s_1).]
Poder optimizar los estados directamente también ayuda con la exploración, ya que podemos navegar temporalmente a través de dominios no físicos para encontrar el plan óptimo:
Sin embargo, el almuerzo nunca es gratis. Y, de hecho, especialmente para los modelos mundiales basados en el aprendizaje profundo, existe un problema crítico que hace que la optimización anterior sea bastante difícil en la práctica.
Un problema para los modelos mundiales basados en el aprendizaje profundo: la sensibilidad de los gradientes de entrada de estado
El tl;dr de esta sección es: optimizar directamente los estados a través de un $F_{theta}$ basado en el aprendizaje profundo es increíblemente frágil, al estilo de la robustez adversarial. Incluso si entrena su modelo mundial en un espacio de estados de dimensiones inferiores, el proceso de entrenamiento para el modelo mundial hace que los paisajes de estados invisibles sean muy nítidos, ya sea un estado invisible en sí mismo o simplemente una dirección normal/ortogonal a la variedad de datos.
Robustez adversaria y el modelo de “colector con hoyuelos”
La robustez adversaria originalmente analizó los modelos de clasificación $f_theta : mathbb{R}^{wtimes h times c} to mathbb{R}^K$, y demostró que al seguir el gradiente de un logit particular $nabla f_theta^k$ desde una imagen base $x$ (no de clase $k$), no era necesario avanzar mucho $x' = x + epsilonnabla f_theta^k$ para hacer que $f_theta$ clasifique $x'$ como $k$ (Szegedy et al., 2014; Goodfellow et al., 2015):
Trabajos posteriores han pintado una imagen geométrica de lo que está sucediendo: para datos cercanos a una variedad de baja dimensión $mathcal{M}$, el proceso de entrenamiento controla el comportamiento en direcciones tangenciales, pero no regulariza el comportamiento en direcciones ortogonales, lo que conduce a un comportamiento sensible (Stutz et al., 2019). Dicho de otra manera: $f_theta$ tiene una constante de Lipschitz razonable cuando se consideran solo direcciones tangenciales a la variedad de datos $mathcal{M}$, pero puede tener constantes de Lipschitz muy altas en direcciones normales. De hecho, a menudo beneficia que el modelo sea más nítido en estas direcciones normales, de modo que pueda ajustarse a funciones más complicadas con mayor precisión.
Como resultado, estos ejemplos contradictorios son increíblemente comunes incluso para un solo modelo determinado. Además, esto no es sólo un fenómeno de visión por computadora; También aparecen ejemplos contradictorios en LLM (Wallace et al., 2019) y en RL (Gleave et al., 2019).
Si bien existen métodos para entrenar modelos más robustos adversariamente, existe una compensación conocida entre el rendimiento del modelo y la robustez adversarial (Tsipras et al., 2019): especialmente en presencia de muchas variables débilmente correlacionadas, el modelo debe ser más nítido para lograr un mayor rendimiento. De hecho, la mayoría de los algoritmos de entrenamiento modernos, ya sea en visión por computadora o LLM, no entrenan la robustez del adversario. Por lo tanto, al menos hasta que el aprendizaje profundo vea un cambio importante de régimen, este es un problema con el que estaremos atrapados.
¿Por qué la solidez adversarial es un problema para la planificación del modelo mundial?
Considere un único componente de la pérdida de dinámica que estamos optimizando en el enfoque de estado elevado:
[min_{s_t, a_t, s_{t+1}} |F_theta(s_t, a_t) – s_{t+1}|_2^2]
Centrémonos más solo en el estado base:
[min_{s_t} |F_theta(s_t, a_t) – s_{t+1}|_2^2.]
Dado que los modelos mundiales generalmente se entrenan en trayectorias de estado/acción $(s_1, a_1, s_2, a_2, ldots)$, la variedad de datos de estado para $F_{theta}$ tiene una dimensionalidad limitada por el espacio de acción:
[mathrm{dim}(mathcal{M}_s) le mathrm{dim}(mathcal{A}) + 1 + mathrm{dim}(mathcal{R}),]
donde $mathcal{R}$ es un espacio opcional de aumentos (por ejemplo, traslaciones/rotaciones). Por lo tanto, normalmente podemos esperar que $mathrm{dim}(mathcal{M}_s)$ sea mucho menor que $mathrm{dim}(mathcal{S})$ y, por lo tanto: es muy fácil encontrar ejemplos contradictorios que pirateen cualquier estado a cualquier otro estado deseado.
Como resultado, la optimización de la dinámica.
[sum_{t=0}^{T-1} big|F_theta(s_t,a_t) – s_{t+1}big|_2^2]
se siente increíblemente “pegajoso”, ya que los puntos base $s_t$ pueden engañar fácilmente a $F_{theta}$ haciéndoles creer que ya alcanzó su objetivo local.1
1. Este problema de robustez adversarial, si bien es particularmente malo para los enfoques de estados elevados, no es exclusivo de ellos. Incluso para los métodos de optimización en serie que optimizan a través del mapa de implementación completo $mathcal{F}^T$, es posible llegar a estados invisibles, donde es muy fácil tener un componente normal alimentado en los componentes normales sensibles de $D_s F_{theta}$. La acción de la expansión de la regla de la cadena jacobiana es
[Bigl(prod_{t=1}^T D_s F_theta(s_t, a_t)Bigr) D_{a_0}F_theta(s_0, a_0).]
Vea qué sucede si alguna etapa del producto tiene algún componente normal al colector de datos. ↩
Nuestra solución
Aquí es donde entra en juego nuestro nuevo planificador GRASP. La observación principal: si bien $D_s F_{theta}$ no es confiable y conflictivo, el espacio de acción suele ser de baja dimensión y está entrenado exhaustivamente, por lo que $D_a F_{theta}$ es en realidad razonable para optimizar y no sufre el problema de robustez adversarial.
En esencia, GRASP construye un planificador basado en estados elevados/colocación de primer orden que solo depende de la acción jacobiana a través del modelo mundial. De esta manera explotamos la diferenciabilidad de los modelos mundiales aprendidos $F_{theta}$, sin ser víctimas de la sensibilidad inherente de los jacobianos estatales $D_s F_{theta}$.
GRASP: Planificador estocástico degradado relajado
Como se señaló anteriormente, comenzamos con el objetivo de planificación de colocación, donde elevamos los estados y relajamos la dinámica hasta convertirla en una penalización:
[min_{mathbf{s},mathbf{a}} mathcal{L}(mathbf{s}, mathbf{a}) = sum_{t=0}^{T-1} big|F_theta(s_t,a_t) – s_{t+1}big|_2^2, quad text{con } s_0 text{ fijo y } s_T=g.]
Luego hacemos dos adiciones clave.
Ingrediente 1: Exploración haciendo ruido en las iteraciones del estado.
Incluso con un objetivo más fluido, la planificación no es convexa. Introducimos la exploración inyectando ruido gaussiano en las actualizaciones del estado virtual durante la optimización.
Una versión sencilla:
[s_t leftarrow s_t – eta_s nabla_{s_t}mathcal{L} + sigma_{text{state}} xi, qquad xisimmathcal{N}(0,I).]
Las acciones todavía se actualizan mediante descenso no estocástico:
[a_t leftarrow a_t – eta_a nabla_{a_t}mathcal{L}.]
El ruido de estado le ayuda a "saltar" entre cuencas en el espacio elevado, mientras que las acciones siguen guiadas por gradientes. Descubrimos que aquí los estados específicamente ruidosos (a diferencia de las acciones) encuentran un buen equilibrio entre la exploración y la capacidad de encontrar mínimos más nítidos.2
2. Debido a que sólo analizamos los estados (y no las acciones), las dinámicas correspondientes no son verdaderamente dinámicas de Langevin. ↩
Ingrediente 2: Reformar gradientes: detener gradientes de entrada de estado frágiles, mantener gradientes de acción
Como se analizó, la vía frágil es el gradiente que fluye hacia la entrada de estado del modelo mundial, (D_s F_{theta}) . La forma más sencilla de hacer esto inicialmente es simplemente detener los gradientes de estado en (F_{theta}) directamente:
Sea $bar{s}_t$ el mismo valor que $s_t$, pero con los gradientes detenidos.
Defina la pérdida de dinámica del gradiente de parada:
[mathcal{L}_{text{dyn}}^{text{sg}}(mathbf{s},mathbf{a}) = sum_{t=0}^{T-1} big|F_theta(bar{s}_t, a_t) – s_{t+1}big|_2^2.]
Esto por sí solo no funciona. Observe que ahora los estados solo siguen el paso del estado anterior, sin que nada obligue a los estados base a perseguir los siguientes. Como resultado, existen mínimos triviales para simplemente detenerse en el origen y luego solo para la acción final que intenta llegar a la meta en un solo paso.
Configuración densa de objetivos
Podemos ver el problema anterior como que la señal del objetivo está completamente aislada de los estados anteriores. Una forma de solucionar este problema es simplemente agregar un término objetivo denso a lo largo de la predicción:
[mathcal{L}_{text{meta}}^{text{sg}}(mathbf{s},mathbf{a}) = sum_{t=0}^{T-1} big|F_theta(bar{s}_t, a_t) – gbig|_2^2.]
En entornos normales, esto provocaría un sesgo excesivo hacia la solución codiciosa de perseguir directamente el objetivo, pero en nuestro entorno esto se equilibra con el sesgo de la pérdida de dinámica de gradiente de parada hacia una dinámica factible. El objetivo final entonces es el siguiente:
[mathcal{L}(mathbf{s},mathbf{a}) = mathcal{L}_{text{dyn}}^{text{sg}}(mathbf{s},mathbf{a}) + gamma , mathcal{L}_{text{objetivo}}^{text{sg}}(mathbf{s},mathbf{a}).]
El resultado es un objetivo de optimización de la planificación que no depende de los gradientes de estado.
“Sincronización” periódica: regrese brevemente a los gradientes de implementación reales
El objetivo de gradiente de parada elevado es excelente para una exploración rápida y guiada, pero sigue siendo una aproximación del objetivo de despliegue en serie original.
Entonces, cada $K_{text{sync}}$ iteraciones, GRASP realiza una breve fase de refinamiento:
Implemente desde $s_0$ usando las acciones actuales $mathbf{a}$ y realice algunos pequeños pasos de gradiente en la pérdida en serie original: [mathbf{a} leftarrow mathbf{a} – eta_{text{sync}},nabla_{mathbf{a}},|s_T(mathbf{a})-g|_2^2.]
La optimización del estado elevado todavía proporciona el núcleo de la optimización, mientras que este paso de refinamiento agrega algo de ayuda para mantener los estados y las acciones basados en trayectorias reales. Por supuesto, este paso de perfeccionamiento puede sustituirse por un planificador en serie de su elección (por ejemplo, CEM); La idea central es seguir obteniendo algunos de los beneficios de la sincronización de ruta completa de los planificadores en serie, y al mismo tiempo seguir utilizando principalmente los beneficios de la planificación de estado elevado.
Cómo GRASP aborda la planificación a largo plazo
Los planificadores basados en la colocación ofrecen una solución natural para la planificación a largo plazo, pero esta optimización es bastante difícil en los modelos del mundo moderno debido a problemas de robustez contradictorios. GRASP propone una solución simple para un planificador basado en la colocación más fluido, junto con una estocasticidad estable para la exploración. Como resultado, la planificación a más largo plazo termina no sólo teniendo más éxito, sino también logrando dichos éxitos más rápidamente:
Resultados Push-T. Tasa de éxito (%)/tiempo medio hasta el éxito. Negrita = mejor de la fila. Tenga en cuenta que el tiempo medio de éxito será mayor cuanto mayor sea la tasa de éxito; GRASP logra ser más rápido a pesar de una mayor tasa de éxito.
¿Qué sigue?
Todavía queda mucho trabajo por hacer para los planificadores de modelos del mundo moderno. Queremos explotar la estructura de gradiente de los modelos mundiales aprendidos, y la colocación (optimización de estado elevado) es un enfoque natural para la planificación a largo plazo, pero es crucial comprender la estructura de gradiente típica aquí: gradientes de acción suaves e informativos y gradientes de estado frágiles. Consideramos GRASP como una iteración inicial para dichos planificadores.
La extensión a modelos mundiales basados en la difusión (los pasos de tiempo latentes más profundos pueden verse como versiones suavizadas del propio modelo mundial), optimizadores y estrategias de ruido más sofisticados, y la integración de GRASP en un sistema de circuito cerrado o en el aprendizaje de políticas de RL para una planificación adaptativa a largo plazo son todos los próximos pasos naturales e interesantes.
Realmente creo que es un momento emocionante para trabajar en planificadores de modelos mundiales. Es un punto dulce divertido donde la literatura de fondo (planificación y control en general) es increíblemente madura y está bien desarrollada, pero el entorno actual (optimización pura de la planificación sobre modelos mundiales modernos a gran escala) todavía está muy poco explorado. Pero, una vez que descubramos todas las ideas correctas, los planificadores de modelos mundiales probablemente se volverán tan comunes como la vida real.
Para obtener más detalles, lea el artículo completo o visite el sitio web del proyecto.