342 lines
14 KiB
Python
342 lines
14 KiB
Python
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_uses_pseudo_huber_and_passes_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']
|
|
noise = torch.tensor(
|
|
[
|
|
[[0.2, -0.4], [0.1, 0.3], [0.5, -0.2]],
|
|
[[-0.1, 0.25], [0.4, -0.5], [0.2, 0.6]],
|
|
],
|
|
dtype=torch.float32,
|
|
)
|
|
t_sample = torch.full((2,), 0.75, dtype=torch.float32)
|
|
r_sample = torch.full((2,), 0.25, dtype=torch.float32)
|
|
with mock.patch('roboimi.vla.agent_imf.torch.randn_like', return_value=noise), \
|
|
mock.patch(
|
|
'roboimi.vla.agent_imf.torch.rand',
|
|
side_effect=[t_sample, r_sample],
|
|
):
|
|
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])))
|
|
cond = head.calls[0]['cond'].float()
|
|
cond_term = cond.mean(dim=(1, 2), keepdim=True)
|
|
z_t = (1 - t_sample.view(2, 1, 1)) * actions + t_sample.view(2, 1, 1) * noise
|
|
scale = head.scale.detach()
|
|
u = scale * z_t + r_sample.view(2, 1, 1) + 2.0 * t_sample.view(2, 1, 1) + cond_term
|
|
v = scale * z_t + 3.0 * t_sample.view(2, 1, 1) + cond_term
|
|
du_dt = scale * v + 2.0
|
|
compound_velocity = u + (t_sample - r_sample).view(2, 1, 1) * du_dt
|
|
expected_loss = (torch.sqrt(1.0 + (compound_velocity - noise).square()) - 1.0).mean()
|
|
mse_loss = (compound_velocity - noise).square().mean()
|
|
self.assertAlmostEqual(loss.item(), expected_loss.item(), places=4)
|
|
self.assertLess(loss.item(), mse_loss.item())
|
|
|
|
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)
|
|
self.assertEqual(agent.loss_type, 'pseudo_huber')
|
|
self.assertEqual(agent.pseudo_huber_delta, 1.0)
|
|
|
|
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)
|
|
self.assertEqual(cfg.agent.loss_type, 'pseudo_huber')
|
|
self.assertEqual(cfg.agent.pseudo_huber_delta, 1.0)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|