Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion src/mobius/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,13 +22,16 @@
"ModelRegistration",
"ModelRegistry",
"ModelTask",
"MLPWorldModel",
"MMSConfig",
"OPSET_VERSION",
"Sam2Config",
"SegformerConfig",
"VisionConfig",
"VisionLanguageConfig",
"WhisperConfig",
"WorldModelConfig",
"WorldModelTask",
"YolosConfig",
"apply_weights",
"build",
Expand Down Expand Up @@ -76,6 +79,7 @@
VisionConfig,
VisionLanguageConfig,
WhisperConfig,
WorldModelConfig,
YolosConfig,
)
from mobius._constants import OPSET_VERSION
Expand All @@ -91,4 +95,5 @@
from mobius._weight_loading import apply_weights
from mobius.integrations.gguf import build_from_gguf
from mobius.integrations.nemo import build_from_nemo
from mobius.tasks import CausalLMTask, ModelTask
from mobius.models import MLPWorldModel
from mobius.tasks import CausalLMTask, ModelTask, WorldModelTask
2 changes: 2 additions & 0 deletions src/mobius/_configs/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,7 @@
TTSConfig,
VisionConfig,
)
from mobius._configs._world_model import WorldModelConfig

__all__ = [
"DEFAULT_INT",
Expand Down Expand Up @@ -117,6 +118,7 @@
"VisionConfig",
"VisionLanguageConfig",
"WhisperConfig",
"WorldModelConfig",
"YolosConfig",
"Zamba2Config",
"_as_int",
Expand Down
64 changes: 64 additions & 0 deletions src/mobius/_configs/_world_model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT License.

"""Configuration for directly declared world models."""

from __future__ import annotations

import dataclasses
import math

from mobius._configs._base import BaseModelConfig


@dataclasses.dataclass
class WorldModelConfig(BaseModelConfig):
"""Configuration shared by single-step world-model graphs.

The three shapes exclude the leading batch dimension. The default
:class:`~mobius.models.MLPWorldModel` flattens each value internally, while
custom modules may preserve their original ranks.
"""

observation_shape: tuple[int, ...] = (1,)
action_shape: tuple[int, ...] = (1,)
state_shape: tuple[int, ...] = (1,)
hidden_size: int = 128
num_hidden_layers: int = 2
hidden_act: str | None = "silu"
residual_state: bool = True

@property
def observation_size(self) -> int:
"""Flattened observation size."""
return math.prod(self.observation_shape)

@property
def action_size(self) -> int:
"""Flattened action size."""
return math.prod(self.action_shape)

@property
def state_size(self) -> int:
"""Flattened recurrent-state size."""
return math.prod(self.state_shape)

def validate(self) -> None:
"""Validate dimensions required by the world-model task and reference model."""
for name, shape in (
("observation_shape", self.observation_shape),
("action_shape", self.action_shape),
("state_shape", self.state_shape),
):
if not shape:
raise ValueError(f"{name} must contain at least one dimension")
if any(
not isinstance(dim, int) or isinstance(dim, bool) or dim <= 0 for dim in shape
):
raise ValueError(f"{name} must contain only positive integer dimensions")
if self.hidden_size <= 0:
raise ValueError("hidden_size must be positive")
if self.num_hidden_layers <= 0:
raise ValueError("num_hidden_layers must be positive")
if self.hidden_act is None:
raise ValueError("hidden_act must be set")
2 changes: 2 additions & 0 deletions src/mobius/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,6 +149,7 @@
"Wav2Vec2ForCTCModel",
"Wav2Vec2Model",
"WhisperForConditionalGeneration",
"MLPWorldModel",
"XLMCausalLMModel",
"Zamba2CausalLMModel",
"mimi_default_config",
Expand Down Expand Up @@ -309,5 +310,6 @@
from mobius.models.wav2vec2 import Wav2Vec2Model
from mobius.models.wav2vec2_ctc import Wav2Vec2ForCTCModel
from mobius.models.whisper import WhisperForConditionalGeneration
from mobius.models.world_model import MLPWorldModel
from mobius.models.xlm import XLMCausalLMModel
from mobius.models.zamba2 import Zamba2CausalLMModel
87 changes: 87 additions & 0 deletions src/mobius/models/world_model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT License.

"""Minimal directly declared world-model implementation."""

from __future__ import annotations

import torch
from onnxscript import OpBuilder, nn

from mobius._configs import WorldModelConfig
from mobius.components import Linear, get_activation


class MLPWorldModel(nn.Module):
"""Deterministic MLP reference model for the world-model task contract."""

default_task = "world-model"
config_class = WorldModelConfig
category = "World Model"

def __init__(self, config: WorldModelConfig):
super().__init__()
config.validate()
self.config = config
input_size = config.observation_size + config.action_size + config.state_size
self.input_layer = Linear(input_size, config.hidden_size)
self.hidden_layers = nn.ModuleList(
[
Linear(config.hidden_size, config.hidden_size)
for _ in range(config.num_hidden_layers - 1)
]
)
self.state_head = Linear(config.hidden_size, config.state_size)
self.observation_head = Linear(config.hidden_size, config.observation_size)
self.reward_head = Linear(config.hidden_size, 1)
self.continuation_head = Linear(config.hidden_size, 1)
self._activation = get_activation(config.hidden_act)

def forward(self, op: OpBuilder, observation, action, state):
observation_flat = op.Flatten(observation, axis=1)
action_flat = op.Flatten(action, axis=1)
state_flat = op.Flatten(state, axis=1)

hidden = self._activation(
op,
self.input_layer(
op,
op.Concat(observation_flat, action_flat, state_flat, axis=1),
),
)
for layer in self.hidden_layers:
hidden = self._activation(op, layer(op, hidden))

next_state_flat = self.state_head(op, hidden)
if self.config.residual_state:
next_state_flat = op.Add(state_flat, next_state_flat)

observation_prediction_flat = self.observation_head(op, hidden)
reward = self.reward_head(op, hidden)
continuation = op.Sigmoid(self.continuation_head(op, hidden))

next_state = self._reshape_batch(
op,
next_state_flat,
state,
self.config.state_shape,
)
observation_prediction = self._reshape_batch(
op,
observation_prediction_flat,
observation,
self.config.observation_shape,
)
return next_state, observation_prediction, reward, continuation

@staticmethod
def _reshape_batch(op: OpBuilder, value, batch_source, shape: tuple[int, ...]):
batch = op.Shape(batch_source, start=0, end=1)
return op.Reshape(value, op.Concat(batch, list(shape), axis=0))

def preprocess_weights(
self,
state_dict: dict[str, torch.Tensor],
) -> dict[str, torch.Tensor]:
"""Return weights unchanged; provided for parity with other Mobius models."""
return state_dict
Loading
Loading