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()