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:
Logic
2026-05-23 22:35:47 +08:00
parent acbd7c605a
commit d94eb8f70b
15 changed files with 2069 additions and 27 deletions
+227
View File
@@ -14,6 +14,8 @@ from roboimi.demos.vla_scripts import eval_vla, train_vla
class _FakeDataset:
available_episode_indices = [0, 1]
def __len__(self):
return 4
@@ -29,6 +31,13 @@ class _FakeLoader:
return iter(self._batches)
class _FakeValDataset(_FakeDataset):
available_episode_indices = [1]
def __len__(self):
return 2
class _FakeOptimizer:
def __init__(self, lr=1e-3):
self.param_groups = [{'lr': lr}]
@@ -91,6 +100,16 @@ class _FakeAgent(nn.Module):
return {}
class _CapturingAgent(_FakeAgent):
def __init__(self):
super().__init__()
self.compute_loss_inputs = []
def compute_loss(self, agent_input):
self.compute_loss_inputs.append(agent_input)
return (self.weight - torch.tensor(0.5)).pow(2)
class _SequentialLossAgent(nn.Module):
def __init__(self, losses):
super().__init__()
@@ -150,6 +169,94 @@ class _FakeEvalEnv:
class TrainVLARolloutValidationTest(unittest.TestCase):
def test_run_training_passes_variable_batch_task_to_agent_input(self):
cfg = OmegaConf.create(
{
'train': {
'device': 'cpu',
'batch_size': 2,
'num_workers': 0,
'val_split': 0.0,
'seed': 0,
'lr': 1e-3,
'max_steps': 1,
'log_freq': 100,
'save_freq': 1000,
'warmup_steps': 1,
'scheduler_type': 'constant',
'min_lr': 0.0,
'grad_clip': 1.0,
'weight_decay': 0.0,
'pretrained_ckpt': None,
'resume_ckpt': None,
'use_swanlab': False,
'rollout_val_freq_epochs': 0,
'rollout_validate_on_checkpoint': False,
'rollout_num_episodes': 1,
},
'data': {
'camera_names': ['front'],
'dataset_dir': 'unused',
},
'agent': {
'_target_': 'fake.agent',
'normalization_type': 'min_max',
},
'eval': {
'ckpt_path': 'unused.pt',
'num_episodes': 1,
'max_timesteps': 1,
'device': 'cpu',
'task_name': 'sim_transfer',
'camera_names': ['front'],
'use_smoothing': False,
'smooth_alpha': 0.3,
'verbose_action': False,
'headless': True,
},
'experiment': {},
}
)
agent = _CapturingAgent()
batch_task = ['pick the red cube', 'insert the peg into the socket']
def fake_instantiate(config_node, **_kwargs):
if config_node is cfg.data:
return _FakeDataset()
if config_node is cfg.agent:
return agent
raise AssertionError(f'unexpected instantiate config: {config_node!r}')
def fake_dataloader(_dataset, *, shuffle, **_kwargs):
del shuffle, _kwargs
return _FakeLoader(
{
'observation.front': torch.zeros(2, 2, 3, 4, 4),
'observation.state': torch.zeros(2, 2, 4),
'action': torch.zeros(2, 8, 2),
'action_is_pad': torch.zeros(2, 8, dtype=torch.bool),
'task': list(batch_task),
},
length=1,
)
with tempfile.TemporaryDirectory() as tempdir:
previous_cwd = os.getcwd()
try:
os.chdir(tempdir)
with mock.patch.object(train_vla, 'instantiate', side_effect=fake_instantiate), \
mock.patch.object(train_vla, 'DataLoader', side_effect=fake_dataloader), \
mock.patch.object(train_vla, 'build_training_optimizer', return_value=_FakeOptimizer(cfg.train.lr)), \
mock.patch.object(train_vla, 'get_lr_schedule_with_warmup', return_value=_FakeScheduler()), \
mock.patch.object(train_vla, 'tqdm', side_effect=lambda iterable, **kwargs: _FakeProgressBar(iterable)), \
mock.patch.object(train_vla.torch, 'save', return_value=None):
train_vla._run_training(cfg)
finally:
os.chdir(previous_cwd)
self.assertEqual(len(agent.compute_loss_inputs), 1)
self.assertEqual(agent.compute_loss_inputs[0]['task'], batch_task)
def test_default_train_config_uses_full_dataset_and_epoch_rollout_validation(self):
cfg = OmegaConf.load(Path('roboimi/vla/conf/config.yaml'))
@@ -162,6 +269,39 @@ class TrainVLARolloutValidationTest(unittest.TestCase):
self.assertIsNone(cfg.train.rollout_num_workers)
self.assertIsNone(cfg.train.rollout_cuda_devices)
def test_explicit_val_episode_indices_builds_held_out_dataset(self):
cfg = OmegaConf.create(
{
'train': {
'val_episode_indices': [1],
'val_split': 0.0,
'seed': 42,
},
'data': {},
}
)
instantiate_calls = []
def fake_instantiate(config_node, **kwargs):
del config_node
instantiate_calls.append(dict(kwargs))
if kwargs.get('episode_indices') == [1]:
return _FakeValDataset()
return _FakeDataset()
with mock.patch.object(train_vla, 'instantiate', side_effect=fake_instantiate):
dataset, train_dataset, val_dataset, explicit = train_vla.build_train_val_datasets(
cfg,
dataset_image_resize_shape=None,
)
self.assertIsInstance(dataset, _FakeDataset)
self.assertIsInstance(train_dataset, _FakeDataset)
self.assertIsInstance(val_dataset, _FakeValDataset)
self.assertEqual(explicit, [1])
self.assertEqual(instantiate_calls[1]['episode_indices'], [0])
self.assertEqual(instantiate_calls[2]['episode_indices'], [1])
def test_run_training_rollout_validation_propagates_gpu_parallel_settings(self):
cfg = OmegaConf.create(
@@ -340,6 +480,93 @@ class TrainVLARolloutValidationTest(unittest.TestCase):
self.assertIn('image_resize_shape', captured_dataset_kwargs)
self.assertIsNone(captured_dataset_kwargs['image_resize_shape'])
def test_training_passes_condition_encoder_image_resize_override_to_dataset_instantiation(self):
cfg = OmegaConf.create(
{
'agent': {
'condition_encoder': {
'dataset_image_resize_shape': None,
},
'normalization_type': 'min_max',
},
'data': {
'dataset_dir': 'unused',
'camera_names': ['front'],
'image_resize_shape': [224, 224],
},
'train': {
'batch_size': 2,
'lr': 1e-4,
'max_steps': 0,
'device': 'cpu',
'disable_cudnn': False,
'num_workers': 0,
'val_split': 0.0,
'seed': 42,
'log_freq': 1,
'save_freq': 10,
'use_swanlab': False,
'rollout_val_freq_epochs': 0,
'rollout_validate_on_checkpoint': False,
'rollout_num_episodes': 1,
'warmup_steps': 1,
'scheduler_type': 'constant',
'min_lr': 1e-6,
'weight_decay': 1e-5,
'grad_clip': 1.0,
'pretrained_ckpt': None,
},
'eval': {
'ckpt_path': 'unused.pt',
'num_episodes': 1,
'headless': True,
'device': 'cpu',
'verbose_action': False,
},
'experiment': {},
}
)
captured_dataset_kwargs = {}
def fake_instantiate(config_node, **kwargs):
if config_node is cfg.data:
captured_dataset_kwargs.update(kwargs)
return _FakeDataset()
if config_node is cfg.agent:
return _FakeAgent()
raise AssertionError(f'unexpected instantiate config: {config_node!r}')
def fake_dataloader(_dataset, *, shuffle, **_kwargs):
del shuffle, _kwargs
return _FakeLoader(
{
'observation.front': torch.zeros(1, 3, 2, 2),
'observation.state': torch.zeros(1, 4),
'action': torch.zeros(1, 2),
'action_is_pad': torch.zeros(1, 1, dtype=torch.bool),
},
length=1,
)
with tempfile.TemporaryDirectory() as tempdir:
previous_cwd = os.getcwd()
try:
os.chdir(tempdir)
with mock.patch.object(train_vla, 'instantiate', side_effect=fake_instantiate), \
mock.patch.object(train_vla, 'DataLoader', side_effect=fake_dataloader), \
mock.patch.object(train_vla, 'build_training_optimizer', return_value=_FakeOptimizer(cfg.train.lr)), \
mock.patch.object(train_vla, 'get_lr_schedule_with_warmup', return_value=_FakeScheduler()), \
mock.patch.object(train_vla, 'tqdm', side_effect=lambda iterable, **kwargs: _FakeProgressBar(iterable)), \
mock.patch.object(train_vla, '_init_swanlab', return_value=None), \
mock.patch.object(train_vla, '_finish_swanlab', return_value=None), \
mock.patch.object(train_vla.torch, 'save', return_value=None):
train_vla._run_training(cfg)
finally:
os.chdir(previous_cwd)
self.assertIn('image_resize_shape', captured_dataset_kwargs)
self.assertIsNone(captured_dataset_kwargs['image_resize_shape'])
def test_eval_main_delegates_to_plain_run_eval_helper(self):
cfg = OmegaConf.create(
{