feat(vla): add SmolVLA conditioned agent with prefix encoder and train/eval validation
Introduce SmolVLA-style VLM prefix encoder backbone and IMF-AttnRes conditioned agent. Add episode-level train/val split, action MSE validation in training, and headless eval support.
This commit is contained in:
@@ -0,0 +1,311 @@
|
||||
import contextlib
|
||||
import importlib
|
||||
import importlib.machinery
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest import mock
|
||||
|
||||
import torch
|
||||
from hydra import compose, initialize_config_dir
|
||||
from hydra.core.global_hydra import GlobalHydra
|
||||
from hydra.utils import instantiate
|
||||
from omegaconf import OmegaConf
|
||||
from torch import nn
|
||||
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
_CONFIG_DIR = str((_REPO_ROOT / 'roboimi/vla/conf').resolve())
|
||||
_CAMERA_NAMES = ('r_vis', 'top', 'front')
|
||||
_MISSING = object()
|
||||
|
||||
|
||||
class _RecordingHead(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.scale = nn.Parameter(torch.tensor(0.5))
|
||||
self.calls = []
|
||||
|
||||
@staticmethod
|
||||
def _broadcast(value, reference):
|
||||
while value.ndim < reference.ndim:
|
||||
value = value.unsqueeze(-1)
|
||||
return value
|
||||
|
||||
def forward(self, sample, r, t, cond=None):
|
||||
self.calls.append({
|
||||
'sample': sample.detach().clone(),
|
||||
'r': r.detach().clone(),
|
||||
't': t.detach().clone(),
|
||||
'cond': None if cond is None else cond.detach().clone(),
|
||||
})
|
||||
cond_term = 0.0 if cond is None else cond.mean(dim=(1, 2), keepdim=True)
|
||||
return self.scale * sample + self._broadcast(r, sample) + 2.0 * self._broadcast(t, sample) + cond_term
|
||||
|
||||
|
||||
class _TaskAwareConditionEncoder(nn.Module):
|
||||
output_dim = 4
|
||||
condition_sequence_length = 3
|
||||
tokens_per_step = 3
|
||||
joint_output_dim = 4
|
||||
camera_names = _CAMERA_NAMES
|
||||
num_cameras = 3
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__()
|
||||
self.constructor_kwargs = dict(kwargs)
|
||||
self.bias = nn.Parameter(torch.tensor(0.0))
|
||||
self.calls = []
|
||||
|
||||
def forward(self, images, state, task=None):
|
||||
self.calls.append({'task': task, 'state': state.detach().clone(), 'image_keys': tuple(images.keys())})
|
||||
batch_size = state.shape[0]
|
||||
state_last = state[:, -1, 0]
|
||||
if isinstance(task, str):
|
||||
task_lengths = torch.full((batch_size,), float(len(task)), dtype=state.dtype, device=state.device)
|
||||
else:
|
||||
task_lengths = torch.tensor([float(len(item)) for item in task], dtype=state.dtype, device=state.device)
|
||||
image_marker = images['r_vis'][:, -1].mean(dim=(1, 2, 3))
|
||||
token0 = torch.stack([state_last, task_lengths, image_marker, torch.ones_like(state_last)], dim=-1)
|
||||
token1 = token0 + 1.0
|
||||
token2 = token0 + 2.0
|
||||
return torch.stack([token0, token1, token2], dim=1) + self.bias
|
||||
|
||||
|
||||
class _BF16TaskAwareConditionEncoder(_TaskAwareConditionEncoder):
|
||||
def forward(self, images, state, task=None):
|
||||
return super().forward(images, state, task=task).to(dtype=torch.bfloat16)
|
||||
|
||||
|
||||
class _StubIMFHead(nn.Module):
|
||||
def __init__(self, input_dim, output_dim, horizon, n_obs_steps, cond_dim, **kwargs):
|
||||
super().__init__()
|
||||
self.constructor_kwargs = {
|
||||
'input_dim': input_dim,
|
||||
'output_dim': output_dim,
|
||||
'horizon': horizon,
|
||||
'n_obs_steps': n_obs_steps,
|
||||
'cond_dim': cond_dim,
|
||||
**kwargs,
|
||||
}
|
||||
self.proj = nn.Linear(input_dim, output_dim)
|
||||
self.cond_obs_emb = nn.Linear(cond_dim, max(cond_dim, 1))
|
||||
|
||||
def forward(self, sample, r, t, cond=None):
|
||||
return torch.zeros_like(sample)
|
||||
|
||||
def get_optim_groups(self, weight_decay):
|
||||
return [
|
||||
{'params': [self.proj.weight], 'weight_decay': weight_decay},
|
||||
{'params': [self.proj.bias, self.cond_obs_emb.weight, self.cond_obs_emb.bias], 'weight_decay': 0.0},
|
||||
]
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _stub_optional_modules(include_head=False, include_condition_encoder=False):
|
||||
previous = {}
|
||||
|
||||
def inject(name, module):
|
||||
if name not in previous:
|
||||
previous[name] = sys.modules.get(name, _MISSING)
|
||||
sys.modules[name] = module
|
||||
|
||||
diffusers_module = types.ModuleType('diffusers')
|
||||
schedulers_module = types.ModuleType('diffusers.schedulers')
|
||||
ddpm_module = types.ModuleType('diffusers.schedulers.scheduling_ddpm')
|
||||
ddim_module = types.ModuleType('diffusers.schedulers.scheduling_ddim')
|
||||
|
||||
class _FakeScheduler:
|
||||
def __init__(self, num_train_timesteps=100, **kwargs):
|
||||
self.config = types.SimpleNamespace(num_train_timesteps=num_train_timesteps)
|
||||
|
||||
ddpm_module.DDPMScheduler = _FakeScheduler
|
||||
ddim_module.DDIMScheduler = _FakeScheduler
|
||||
diffusers_module.DDPMScheduler = _FakeScheduler
|
||||
diffusers_module.DDIMScheduler = _FakeScheduler
|
||||
diffusers_module.schedulers = schedulers_module
|
||||
|
||||
try:
|
||||
inject('diffusers', diffusers_module)
|
||||
inject('diffusers.schedulers', schedulers_module)
|
||||
inject('diffusers.schedulers.scheduling_ddpm', ddpm_module)
|
||||
inject('diffusers.schedulers.scheduling_ddim', ddim_module)
|
||||
if include_head:
|
||||
import roboimi.vla.models.heads as heads_package
|
||||
head_module = types.ModuleType('roboimi.vla.models.heads.imf_transformer1d')
|
||||
head_module.IMFTransformer1D = _StubIMFHead
|
||||
inject('roboimi.vla.models.heads.imf_transformer1d', head_module)
|
||||
setattr(heads_package, 'imf_transformer1d', head_module)
|
||||
if include_condition_encoder:
|
||||
module = types.ModuleType('tests.fake_smolvla_condition_encoder')
|
||||
module.TaskAwareConditionEncoder = _TaskAwareConditionEncoder
|
||||
module.BF16TaskAwareConditionEncoder = _BF16TaskAwareConditionEncoder
|
||||
inject('tests.fake_smolvla_condition_encoder', module)
|
||||
yield
|
||||
finally:
|
||||
for name, old in reversed(list(previous.items())):
|
||||
if old is _MISSING:
|
||||
sys.modules.pop(name, None)
|
||||
else:
|
||||
sys.modules[name] = old
|
||||
|
||||
|
||||
def _compose_cfg(overrides=None):
|
||||
if not OmegaConf.has_resolver('len'):
|
||||
OmegaConf.register_new_resolver('len', lambda x: len(x))
|
||||
GlobalHydra.instance().clear()
|
||||
with initialize_config_dir(version_base=None, config_dir=_CONFIG_DIR):
|
||||
return compose(config_name='config', overrides=list(overrides or []))
|
||||
|
||||
|
||||
class SmolVLAIMFAgentTest(unittest.TestCase):
|
||||
def test_compute_loss_and_predict_action_pass_variable_task_to_condition_encoder(self):
|
||||
from roboimi.vla.agent_smolvla_conditioned import SmolVLAIMFAttnResAgent
|
||||
|
||||
condition_encoder = _BF16TaskAwareConditionEncoder()
|
||||
head = _RecordingHead()
|
||||
agent = SmolVLAIMFAttnResAgent(
|
||||
condition_encoder=condition_encoder,
|
||||
action_encoder=nn.Identity(),
|
||||
head=head,
|
||||
action_dim=2,
|
||||
obs_dim=1,
|
||||
pred_horizon=3,
|
||||
obs_horizon=2,
|
||||
diffusion_steps=10,
|
||||
inference_steps=1,
|
||||
num_cams=3,
|
||||
camera_names=_CAMERA_NAMES,
|
||||
num_action_steps=2,
|
||||
head_type='transformer',
|
||||
)
|
||||
images = {
|
||||
'r_vis': torch.full((2, 2, 1, 2, 2), 1.0),
|
||||
'top': torch.full((2, 2, 1, 2, 2), 2.0),
|
||||
'front': torch.full((2, 2, 1, 2, 2), 3.0),
|
||||
}
|
||||
qpos = torch.tensor([[[1.0], [2.0]], [[3.0], [4.0]]])
|
||||
actions = torch.zeros(2, 3, 2)
|
||||
tasks = ['short', 'longer task']
|
||||
loss = agent.compute_loss({'images': images, 'qpos': qpos, 'action': actions, 'task': tasks})
|
||||
self.assertTrue(torch.isfinite(loss))
|
||||
self.assertEqual(condition_encoder.calls[-1]['task'], tasks)
|
||||
self.assertEqual(head.calls[-1]['cond'].shape, (2, 3, 4))
|
||||
self.assertTrue(torch.allclose(head.calls[-1]['cond'][:, 0, 1], torch.tensor([5.0, 11.0])))
|
||||
|
||||
with mock.patch('roboimi.vla.agent_imf.torch.randn', return_value=torch.zeros(2, 3, 2)):
|
||||
pred = agent.predict_action(images, qpos, task=tasks)
|
||||
self.assertEqual(pred.shape, (2, 3, 2))
|
||||
self.assertEqual(condition_encoder.calls[-1]['task'], tasks)
|
||||
|
||||
def test_condition_tokens_are_cast_to_action_head_dtype(self):
|
||||
from roboimi.vla.agent_smolvla_conditioned import SmolVLAIMFAttnResAgent
|
||||
|
||||
condition_encoder = _TaskAwareConditionEncoder()
|
||||
head = _RecordingHead()
|
||||
agent = SmolVLAIMFAttnResAgent(
|
||||
condition_encoder=condition_encoder,
|
||||
action_encoder=nn.Identity(),
|
||||
head=head,
|
||||
action_dim=2,
|
||||
obs_dim=1,
|
||||
pred_horizon=3,
|
||||
obs_horizon=2,
|
||||
diffusion_steps=10,
|
||||
inference_steps=1,
|
||||
num_cams=3,
|
||||
camera_names=_CAMERA_NAMES,
|
||||
num_action_steps=2,
|
||||
head_type='transformer',
|
||||
)
|
||||
images = {
|
||||
'r_vis': torch.full((1, 2, 1, 2, 2), 1.0),
|
||||
'top': torch.full((1, 2, 1, 2, 2), 2.0),
|
||||
'front': torch.full((1, 2, 1, 2, 2), 3.0),
|
||||
}
|
||||
qpos = torch.tensor([[[1.0], [2.0]]], dtype=torch.float32)
|
||||
|
||||
cond = agent._build_cond(images, qpos, task=['pick'])
|
||||
|
||||
self.assertEqual(cond.dtype, head.scale.dtype)
|
||||
|
||||
def test_unknown_dataset_task_uses_configured_task_description(self):
|
||||
from roboimi.vla.agent_smolvla_conditioned import SmolVLAIMFAttnResAgent
|
||||
|
||||
condition_encoder = _TaskAwareConditionEncoder()
|
||||
head = _RecordingHead()
|
||||
agent = SmolVLAIMFAttnResAgent(
|
||||
condition_encoder=condition_encoder,
|
||||
action_encoder=nn.Identity(),
|
||||
head=head,
|
||||
action_dim=2,
|
||||
obs_dim=1,
|
||||
pred_horizon=3,
|
||||
obs_horizon=2,
|
||||
diffusion_steps=10,
|
||||
inference_steps=1,
|
||||
num_cams=3,
|
||||
camera_names=_CAMERA_NAMES,
|
||||
num_action_steps=2,
|
||||
head_type='transformer',
|
||||
task_description='insert the peg into the socket',
|
||||
)
|
||||
images = {
|
||||
'r_vis': torch.full((2, 2, 1, 2, 2), 1.0),
|
||||
'top': torch.full((2, 2, 1, 2, 2), 2.0),
|
||||
'front': torch.full((2, 2, 1, 2, 2), 3.0),
|
||||
}
|
||||
qpos = torch.tensor([[[1.0], [2.0]], [[3.0], [4.0]]], dtype=torch.float32)
|
||||
|
||||
agent._build_cond(images, qpos, task=['unknown', ''])
|
||||
|
||||
self.assertEqual(
|
||||
condition_encoder.calls[-1]['task'],
|
||||
['insert the peg into the socket', 'insert the peg into the socket'],
|
||||
)
|
||||
|
||||
def test_hydra_config_instantiates_smolvla_imf_attnres_with_condition_encoder_contract(self):
|
||||
cfg = _compose_cfg(overrides=[
|
||||
'agent=smolvla_imf_attnres',
|
||||
'agent.condition_encoder._target_=tests.fake_smolvla_condition_encoder.TaskAwareConditionEncoder',
|
||||
'agent.condition_dim=4',
|
||||
'agent.condition_sequence_length=3',
|
||||
'agent.head.n_layer=1',
|
||||
'agent.head.n_emb=16',
|
||||
])
|
||||
self.assertEqual(cfg.agent._target_, 'roboimi.vla.agent_smolvla_conditioned.SmolVLAIMFAttnResAgent')
|
||||
self.assertEqual(cfg.agent.head.cond_dim, cfg.agent.condition_dim)
|
||||
self.assertEqual(cfg.agent.head.n_obs_steps, cfg.agent.condition_sequence_length)
|
||||
|
||||
with _stub_optional_modules(include_head=True, include_condition_encoder=True):
|
||||
agent = instantiate(cfg.agent)
|
||||
|
||||
self.assertEqual(agent.per_step_cond_dim, 4)
|
||||
self.assertEqual(agent.condition_sequence_length, 3)
|
||||
self.assertIsInstance(agent.noise_pred_net, _StubIMFHead)
|
||||
self.assertEqual(agent.noise_pred_net.constructor_kwargs['cond_dim'], 4)
|
||||
self.assertEqual(agent.noise_pred_net.constructor_kwargs['n_obs_steps'], 3)
|
||||
|
||||
def test_hydra_config_exposes_smolvla_pretrained_vlm_defaults(self):
|
||||
cfg = _compose_cfg(overrides=[
|
||||
'agent=smolvla_imf_attnres',
|
||||
])
|
||||
|
||||
self.assertEqual(cfg.agent.condition_encoder.model_name, 'HuggingFaceTB/SmolVLM2-500M-Video-Instruct')
|
||||
self.assertTrue(cfg.agent.condition_encoder.load_vlm_weights)
|
||||
self.assertEqual(cfg.agent.condition_encoder.num_vlm_layers, 16)
|
||||
self.assertTrue(cfg.agent.condition_encoder.freeze_vlm)
|
||||
self.assertTrue(cfg.agent.condition_encoder.freeze_vision_encoder)
|
||||
self.assertTrue(cfg.agent.condition_encoder.run_text_model)
|
||||
self.assertEqual(cfg.agent.condition_encoder.max_state_dim, 32)
|
||||
self.assertIsNone(cfg.agent.condition_encoder.dataset_image_resize_shape)
|
||||
self.assertIsNone(cfg.agent.condition_encoder.eval_image_resize_shape)
|
||||
self.assertEqual(cfg.agent.condition_dim, 960)
|
||||
self.assertEqual(cfg.agent.condition_sequence_length, 241)
|
||||
self.assertEqual(cfg.agent.head.cond_dim, 960)
|
||||
self.assertEqual(cfg.agent.head.n_obs_steps, 241)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user