Conozca Open Dreamer: una reproducción JAX/Flax del canal modelo Dreamer 4 World, con la receta de entrenamiento completa publicada

Un pequeño grupo de investigadores de IA (Reactor) ha lanzado Open Dreamer, una implementación abierta del modelo mundial Dreamer 4 escrita en JAX y Flax NNX.

Lo que realmente se envió

Se liberaron dos repositorios. next-state/open-dreamer contiene el proceso de capacitación: un tokenizador de video causal, un modelo de dinámica latente condicionada por la acción, generación de implementación y puntuación FVD. reactor-team/open-dreamer tiene un arnés de implementación local mínimo que genera fotogramas a partir de un MP4 y un archivo de acción coincidente.

Un tercer artefacto es la demostración del navegador alojada en el tiempo de ejecución de Reactor. Transmite un mundo de Minecraft generado en tiempo real y expone un interruptor Juego ⟷ Sueño que transfiere la transmisión del juego real al modelo del mundo cuadro por cuadro.

El objetivo declarado era reproducir la investigación del Dreamer 4. El equipo de investigación evitó deliberadamente métodos fuera de ese trabajo de investigación para mantener estrecho el espacio de búsqueda. Comenzaron con CoinRun, un juego de plataformas 2D generado por procedimientos que se puede entrenar en una sola GPU, luego ampliaron el proceso de trabajo a videos de juego estilo Minecraft/VPT.

Arquitectura: una columna vertebral, dos modelos

Tanto el tokenizador como el modelo dinámico utilizan la misma columna vertebral del transformador causal de bloque. Esa columna vertebral alterna dos tipos de atención. Las capas espaciales propagan información entre los elementos de un solo cuadro. Las capas de tiempo causales propagan información entre fotogramas.

El tokenizador es un codificador automático enmascarado basado en transformador en lugar de un VAE. El equipo informa una compresión de aproximadamente 100 × y señala que el diseño no necesita KL ni pérdida adversaria. El enmascaramiento, argumentan, hace que el espacio latente sea más difusible.

El modelo dinámico realiza una predicción del siguiente cuadro y se entrena con modelos de forzamiento de difusión, coincidencia de flujo y atajos. También predice la siguiente acción. En lugar de alternar entre un módulo de transición separado y una política, la implementación se divide en bloques por paso de tiempo de (acción anterior, estado, política). La atención espacial corre dentro de cada bloque; la atención temporal causal conecta bloques a través del tiempo.

Fundamentalmente, los tokens del modelo mundial no pueden leer el token del agente. Por lo tanto, la información sobre tareas y políticas puede influir en los estados futuros sólo a través de la siguiente acción.

La receta de entrenamiento, tal como está configurada.

Las configuraciones de Minecraft enviadas hacen que la receta sea concreta.

El modelo dinámico tiene 1.6B de parámetros: 30 capas causales de bloques, d_model 1920, 30 cabezales de atención y 3 cabezales KV para atención de consultas agrupadas. Cada cuarta capa es una capa de atención del tiempo. Cada paso de tiempo lleva 32 tokens de registro aprendidos y packaging_factor: 2 empaqueta las latentes del tokenizador vecino en cada token espacial dinámico. La atención del tiempo utiliza una ventana corrediza de 192 pasos.

El entrenamiento dura 200.000 pasos con Muon, un programa WSD y una tasa de aprendizaje máxima de 3e-4. Las muestras de acceso directo/arranque se activan en el paso 100.000 en una fracción de lote de 0,25. La caída de la EMA es 0,999.

La configuración del tokenizador emite 512 tokens latentes por cuadro con un ancho de cuello de botella de 16. Los cuadros sin procesar de 360 ​​× 640 se rellenan a 368 × 640, de modo que ambas dimensiones se dividen en parches de 16 × 16. La profundidad del codificador es 12 en d_model 1536; la profundidad del decodificador es 8 en d_model 1024. La probabilidad de enmascaramiento MAE alcanza un máximo de 0,9 y LPIPS se aplica con un peso de 0,2 en la mitad de los pasos de tiempo.

Las acciones de VPT se analizan en 27 canales de acciones binarias más 121 clases categóricas de mouse, sin canales continuos.

Compumaxxing y el muro de la memoria

El equipo de investigación informa una utilización del modelo FLOP del 57% al 58%, frente a un punto de referencia declarado del 60% para un entrenamiento de transformador saludable. El razonamiento es un argumento de línea de techo. En un B200, el cruce entre el ancho de banda y el cálculo se sitúa en 292 FLOP/byte. Alimentar 256 fotogramas por GPU lleva la carga de trabajo más allá de ese punto máximo.

Sharding fue en sentido contrario a las expectativas. Con 1,6 mil millones de parámetros, el estado del modelo (parámetros, gradientes, estado del optimizador y EMA) ocupaba aproximadamente 24 GiB, lo que cabe en un B200. Las activaciones fueron el costo real. El equipo de investigación probó el paralelismo de datos, FSDP, el paralelismo tensorial y el paralelismo de secuencia, y luego se decidió por el paralelismo de datos simple más puntos de control de activación.

La carga de datos se resolvió tokenizando previamente todo el conjunto de datos en archivos .arrayrecord y luego usando Grain con un búfer de captación previa del lado de la GPU. La decodificación de ffmpeg no fue lo suficientemente rápida para mantener alimentadas las GPU.

La sección de estabilidad es la verdadera carga útil.

El equipo de investigación es explícito en que la estabilidad consumía la mayor parte de su tiempo. Su observación clave: la mayoría de los problemas de estabilidad ocurren a pesar de que las pérdidas disminuyen. MSE mejora sin problemas mientras que la calidad de la generación se degrada.

Se documentan seis correcciones. Muon reemplazó a LaProp, que aumentó de manera aleatoria y cada vez más frecuente, en dos ejecuciones de aproximadamente 400 B200 horas cada una. Los pesos de EMA se consideran obligatorios para la inferencia de difusión. La precisión mixta es sensible a los límites: los parámetros permanecen float32, BF16 cubre la mayoría de las activaciones matmul y entradas de atención, y float32 se mantiene para la normalización y el cabezal de salida del flujo dinámico.

En la ponderación de la pérdida, utilizan la predicción x con una pérdida en el espacio v, lo que se reduce a un término de ponderación similar al de Dreamer 4 pero con un denominador al cuadrado. Informan de una mejora pequeña pero notable. El transporte óptimo baricéntrico de minibatch entre ruido y secuencias latentes hizo que la generación de implementación fuera más estable. La parametrización μ se probó y se consideró innecesaria, en parte porque Muon mantiene los hiperparámetros más estables en todos los tamaños de modelos.

Otro resultado de la fase CoinRun: un barrido iso-FLOP colocó el escalado óptimo de cómputo en aproximadamente N∝C0.56 y D∝C0.44.

¿Qué no está en la caja?

El repositorio no incluye el bucle de entrenamiento de RL ni de clonación de comportamiento; un bucle de agente completo de Dreamer 4 BC/RL aparece como un elemento abierto de la hoja de ruta. El trabajo de política de CoinRun descrito en la publicación no se utilizó para Minecraft y no se publicó.

La publicación tampoco publica puntuaciones FVD, aunque scripts/eval_fvd.py se entrega con un arnés basado en I3D configurado para 4 cuadros de contexto y un horizonte de 240 cuadros.

Conclusiones clave

Open Dreamer reproduce el pipeline de Dreamer 4 en JAX/Flax NNX, con código de entrenamiento y una demostración de Minecraft. El modelo dinámico tiene 1.6B de parámetros, 30 capas, d_model 1920, 200K pasos entrenados con Muon. Números de ingeniería informados: 57–58 % de MFU en B200, 256 fotogramas por GPU, ~24 GiB de estado del modelo. El cuello de botella fue la estabilidad, no el rendimiento; las curvas de pérdida ocultaron la mayoría de las regresiones de calidad generacional.

Consulte la publicación del blog y la demostración, el repositorio de capacitación, el repositorio de inferencia y Reactor on X. Todo el crédito por esta investigación es para los investigadores de este proyecto.

Asif Razzaq es el director ejecutivo de Marktechpost Media Inc.. Como emprendedor e ingeniero visionario, Asif está comprometido a aprovechar el potencial de la inteligencia artificial para el bien social. Su esfuerzo más reciente es el lanzamiento de una plataforma de medios de inteligencia artificial, Marktechpost, que se destaca por su cobertura en profundidad del aprendizaje automático y las noticias sobre aprendizaje profundo que es técnicamente sólida y fácilmente comprensible para una amplia audiencia. La plataforma cuenta con más de 2 millones de visitas mensuales, lo que ilustra su popularidad entre el público.