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,328 @@
|
||||
import types
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class _FakeVisionOutput:
|
||||
def __init__(self, last_hidden_state):
|
||||
self.last_hidden_state = last_hidden_state
|
||||
|
||||
|
||||
class _FakeVisionModel(nn.Module):
|
||||
def __init__(self, hidden_size=4):
|
||||
super().__init__()
|
||||
self.dtype = torch.float32
|
||||
self.scale = nn.Parameter(torch.tensor(1.0))
|
||||
self.calls = []
|
||||
self.hidden_size = hidden_size
|
||||
|
||||
def forward(self, pixel_values=None, patch_attention_mask=None):
|
||||
self.calls.append({
|
||||
'pixel_values': pixel_values.detach().clone(),
|
||||
'patch_attention_mask': patch_attention_mask,
|
||||
})
|
||||
pooled = pixel_values.mean(dim=(2, 3)) * self.scale
|
||||
tokens = torch.stack([pooled, pooled + 1.0], dim=1)
|
||||
return _FakeVisionOutput(tokens)
|
||||
|
||||
|
||||
class _FakeConnector(nn.Module):
|
||||
def __init__(self, in_dim=3, out_dim=4):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(in_dim, out_dim, bias=False)
|
||||
with torch.no_grad():
|
||||
self.proj.weight.copy_(
|
||||
torch.tensor(
|
||||
[
|
||||
[1.0, 0.0, 0.0],
|
||||
[0.0, 1.0, 0.0],
|
||||
[0.0, 0.0, 1.0],
|
||||
[1.0, 1.0, 1.0],
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.proj(x)
|
||||
|
||||
|
||||
class _FakeTextModel(nn.Module):
|
||||
def __init__(self, vocab_size=32, hidden_size=4, num_layers=6):
|
||||
super().__init__()
|
||||
self.embed = nn.Embedding(vocab_size, hidden_size)
|
||||
self.layers = nn.ModuleList([nn.Linear(hidden_size, hidden_size) for _ in range(num_layers)])
|
||||
self.norm = nn.Identity()
|
||||
self.forward_calls = []
|
||||
with torch.no_grad():
|
||||
for idx in range(vocab_size):
|
||||
self.embed.weight[idx].fill_(float(idx))
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.embed
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
position_ids=None,
|
||||
past_key_values=None,
|
||||
inputs_embeds=None,
|
||||
use_cache=None,
|
||||
return_dict=True,
|
||||
cache_position=None,
|
||||
**kwargs,
|
||||
):
|
||||
del input_ids, past_key_values, use_cache, cache_position, kwargs
|
||||
self.forward_calls.append({
|
||||
'attention_mask': None if attention_mask is None else attention_mask.detach().clone(),
|
||||
'position_ids': None if position_ids is None else position_ids.detach().clone(),
|
||||
'inputs_embeds': inputs_embeds.detach().clone(),
|
||||
})
|
||||
hidden = inputs_embeds
|
||||
for layer in self.layers:
|
||||
hidden = layer(hidden)
|
||||
hidden = self.norm(hidden)
|
||||
if return_dict:
|
||||
return types.SimpleNamespace(last_hidden_state=hidden)
|
||||
return (hidden,)
|
||||
|
||||
|
||||
class _FakeVLM(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
text_config = types.SimpleNamespace(hidden_size=4, head_dim=2, num_attention_heads=2, num_key_value_heads=1)
|
||||
self.config = types.SimpleNamespace(text_config=text_config)
|
||||
self.model = types.SimpleNamespace(
|
||||
vision_model=_FakeVisionModel(hidden_size=4),
|
||||
connector=_FakeConnector(in_dim=3, out_dim=4),
|
||||
text_model=_FakeTextModel(hidden_size=4, num_layers=6),
|
||||
)
|
||||
|
||||
|
||||
class _FakeTokenizer:
|
||||
fake_image_token_id = 29
|
||||
global_image_token_id = 30
|
||||
|
||||
def __init__(self):
|
||||
self.padding_side = 'left'
|
||||
self.calls = []
|
||||
|
||||
def __call__(self, text, *, padding, max_length, return_tensors, truncation):
|
||||
self.calls.append({
|
||||
'text': list(text),
|
||||
'padding': padding,
|
||||
'max_length': max_length,
|
||||
'return_tensors': return_tensors,
|
||||
'truncation': truncation,
|
||||
'padding_side': self.padding_side,
|
||||
})
|
||||
batch = len(text)
|
||||
ids = torch.zeros(batch, max_length, dtype=torch.long)
|
||||
mask = torch.zeros(batch, max_length, dtype=torch.bool)
|
||||
for row, item in enumerate(text):
|
||||
del item
|
||||
ids[row, :3] = torch.tensor([1, 2, 3])
|
||||
mask[row, :3] = True
|
||||
return {'input_ids': ids, 'attention_mask': mask}
|
||||
|
||||
|
||||
class SmolVLAPrefixEncoderTest(unittest.TestCase):
|
||||
def test_loads_pretrained_vlm_with_local_files_only_crops_layers_and_freezes_vlm(self):
|
||||
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||
|
||||
fake_vlm = _FakeVLM()
|
||||
fake_tokenizer = _FakeTokenizer()
|
||||
with mock.patch(
|
||||
'roboimi.vla.models.backbones.smolvla_prefix_encoder.AutoModelForImageTextToText.from_pretrained',
|
||||
return_value=fake_vlm,
|
||||
) as model_loader, mock.patch(
|
||||
'roboimi.vla.models.backbones.smolvla_prefix_encoder.AutoTokenizer.from_pretrained',
|
||||
return_value=fake_tokenizer,
|
||||
) as tokenizer_loader:
|
||||
encoder = SmolVLAPrefixEncoder(
|
||||
model_name='HuggingFaceTB/SmolVLM2-500M-Video-Instruct',
|
||||
local_files_only=True,
|
||||
num_vlm_layers=2,
|
||||
freeze_vlm=True,
|
||||
max_state_dim=5,
|
||||
tokenizer_max_length=4,
|
||||
camera_names=('r_vis', 'top'),
|
||||
resize_imgs_with_padding=None,
|
||||
)
|
||||
|
||||
model_loader.assert_called_once()
|
||||
self.assertEqual(model_loader.call_args.kwargs['local_files_only'], True)
|
||||
self.assertEqual(model_loader.call_args.kwargs['torch_dtype'], 'bfloat16')
|
||||
tokenizer_loader.assert_called_once_with(
|
||||
'HuggingFaceTB/SmolVLM2-500M-Video-Instruct',
|
||||
local_files_only=True,
|
||||
)
|
||||
self.assertEqual(len(encoder.vlm.model.text_model.layers), 2)
|
||||
self.assertTrue(all(not p.requires_grad for p in encoder.vlm.parameters()))
|
||||
self.assertFalse(encoder.vlm.training)
|
||||
encoder.train()
|
||||
self.assertFalse(encoder.vlm.training)
|
||||
|
||||
def test_accepts_train_eval_resize_compatibility_fields(self):
|
||||
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||
|
||||
encoder = SmolVLAPrefixEncoder(
|
||||
vlm=_FakeVLM(),
|
||||
tokenizer=_FakeTokenizer(),
|
||||
num_vlm_layers=3,
|
||||
camera_names=('r_vis', 'top'),
|
||||
resize_imgs_with_padding=None,
|
||||
dataset_image_resize_shape=None,
|
||||
eval_image_resize_shape=(640, 480),
|
||||
)
|
||||
|
||||
self.assertIsNone(encoder.dataset_image_resize_shape)
|
||||
self.assertEqual(encoder.eval_image_resize_shape, (640, 480))
|
||||
|
||||
def test_embed_prefix_uses_variable_tasks_camera_order_last_state_and_smolvla_masks(self):
|
||||
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||
|
||||
fake_vlm = _FakeVLM()
|
||||
fake_tokenizer = _FakeTokenizer()
|
||||
encoder = SmolVLAPrefixEncoder(
|
||||
vlm=fake_vlm,
|
||||
tokenizer=fake_tokenizer,
|
||||
num_vlm_layers=3,
|
||||
freeze_vlm=True,
|
||||
train_state_proj=True,
|
||||
max_state_dim=5,
|
||||
tokenizer_max_length=4,
|
||||
camera_names=('r_vis', 'top'),
|
||||
resize_imgs_with_padding=None,
|
||||
)
|
||||
with torch.no_grad():
|
||||
encoder.state_proj.weight.zero_()
|
||||
encoder.state_proj.bias.zero_()
|
||||
encoder.state_proj.weight[:, :4] = torch.eye(4)
|
||||
|
||||
images = {
|
||||
'top': torch.full((2, 2, 3, 2, 2), 0.75),
|
||||
'r_vis': torch.full((2, 2, 3, 2, 2), 0.25),
|
||||
}
|
||||
state = torch.tensor(
|
||||
[
|
||||
[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]],
|
||||
[[7.0, 8.0, 9.0], [10.0, 11.0, 12.0]],
|
||||
]
|
||||
)
|
||||
tasks = ['pick red cube', 'open drawer']
|
||||
out = encoder.embed_prefix(images=images, state=state, task=tasks)
|
||||
|
||||
# 2 cameras * 2 image tokens + 4 language tokens + 1 state token
|
||||
self.assertEqual(out.shape, (2, 9, 4))
|
||||
self.assertEqual(encoder.output_dim, 4)
|
||||
self.assertEqual(encoder.tokens_per_step, 9)
|
||||
self.assertEqual(encoder.condition_sequence_length, 9)
|
||||
|
||||
self.assertEqual(fake_tokenizer.calls[-1]['text'], ['pick red cube\n', 'open drawer\n'])
|
||||
self.assertEqual(fake_tokenizer.calls[-1]['padding'], 'max_length')
|
||||
self.assertEqual(fake_tokenizer.calls[-1]['max_length'], 4)
|
||||
self.assertEqual(fake_tokenizer.calls[-1]['padding_side'], 'right')
|
||||
|
||||
# Camera order is r_vis then top, and pixels are mapped [0,1] -> [-1,1].
|
||||
first_camera_pixels = fake_vlm.model.vision_model.calls[0]['pixel_values']
|
||||
second_camera_pixels = fake_vlm.model.vision_model.calls[1]['pixel_values']
|
||||
self.assertTrue(torch.allclose(first_camera_pixels, torch.full((2, 3, 2, 2), -0.5)))
|
||||
self.assertTrue(torch.allclose(second_camera_pixels, torch.full((2, 3, 2, 2), 0.5)))
|
||||
|
||||
# Last token is padded last state projected from [4,5,6,0,0] and [10,11,12,0,0].
|
||||
self.assertTrue(torch.allclose(out[0, -1], torch.tensor([4.0, 5.0, 6.0, 0.0])))
|
||||
self.assertTrue(torch.allclose(out[1, -1], torch.tensor([10.0, 11.0, 12.0, 0.0])))
|
||||
|
||||
pad_mask, att_mask = encoder.last_prefix_pad_mask, encoder.last_prefix_att_mask
|
||||
self.assertEqual(pad_mask.shape, (2, 9))
|
||||
self.assertEqual(att_mask.shape, (2, 9))
|
||||
self.assertTrue(torch.all(att_mask[:, :8] == 0))
|
||||
self.assertTrue(torch.all(att_mask[:, -1] == 1))
|
||||
|
||||
def test_forward_can_run_cropped_frozen_text_model_over_prefix_tokens(self):
|
||||
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||
|
||||
fake_vlm = _FakeVLM()
|
||||
fake_tokenizer = _FakeTokenizer()
|
||||
encoder = SmolVLAPrefixEncoder(
|
||||
vlm=fake_vlm,
|
||||
tokenizer=fake_tokenizer,
|
||||
num_vlm_layers=2,
|
||||
freeze_vlm=True,
|
||||
max_state_dim=5,
|
||||
tokenizer_max_length=4,
|
||||
camera_names=('r_vis',),
|
||||
resize_imgs_with_padding=None,
|
||||
run_text_model=True,
|
||||
)
|
||||
images = {
|
||||
'r_vis': torch.full((2, 1, 3, 2, 2), 0.25),
|
||||
}
|
||||
state = torch.tensor([[[1.0, 2.0, 3.0]], [[4.0, 5.0, 6.0]]])
|
||||
|
||||
out = encoder(images, state=state, task=['pick', 'place'])
|
||||
|
||||
# 1 camera * 2 image tokens + 4 language tokens + 1 state token.
|
||||
self.assertEqual(out.shape, (2, 7, 4))
|
||||
self.assertEqual(len(fake_vlm.model.text_model.layers), 2)
|
||||
self.assertEqual(len(fake_vlm.model.text_model.forward_calls), 1)
|
||||
text_call = fake_vlm.model.text_model.forward_calls[-1]
|
||||
self.assertEqual(tuple(text_call['attention_mask'].shape), (2, 1, 7, 7))
|
||||
self.assertEqual(tuple(text_call['position_ids'].shape), (2, 7))
|
||||
self.assertTrue(torch.all(out != 0))
|
||||
|
||||
def test_frozen_vlm_still_backpropagates_through_text_model_to_state_projection(self):
|
||||
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||
|
||||
fake_vlm = _FakeVLM()
|
||||
encoder = SmolVLAPrefixEncoder(
|
||||
vlm=fake_vlm,
|
||||
tokenizer=_FakeTokenizer(),
|
||||
num_vlm_layers=3,
|
||||
freeze_vlm=True,
|
||||
train_state_proj=True,
|
||||
max_state_dim=5,
|
||||
tokenizer_max_length=4,
|
||||
camera_names=('r_vis',),
|
||||
resize_imgs_with_padding=None,
|
||||
run_text_model=True,
|
||||
)
|
||||
images = {
|
||||
'r_vis': torch.full((2, 1, 3, 2, 2), 0.25),
|
||||
}
|
||||
state = torch.tensor([[[1.0, 2.0, 3.0]], [[4.0, 5.0, 6.0]]])
|
||||
|
||||
out = encoder(images, state=state, task=['pick', 'place'])
|
||||
out[:, -1].sum().backward()
|
||||
|
||||
self.assertIsNotNone(encoder.state_proj.weight.grad)
|
||||
self.assertGreater(float(encoder.state_proj.weight.grad.abs().sum()), 0.0)
|
||||
self.assertTrue(all(param.grad is None for param in fake_vlm.parameters()))
|
||||
|
||||
def test_embed_prefix_rejects_missing_camera_and_task_batch_mismatch(self):
|
||||
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||
|
||||
encoder = SmolVLAPrefixEncoder(
|
||||
vlm=_FakeVLM(),
|
||||
tokenizer=_FakeTokenizer(),
|
||||
num_vlm_layers=3,
|
||||
camera_names=('r_vis', 'top'),
|
||||
resize_imgs_with_padding=None,
|
||||
max_state_dim=5,
|
||||
tokenizer_max_length=4,
|
||||
)
|
||||
images = {'r_vis': torch.rand(2, 1, 3, 2, 2)}
|
||||
state = torch.rand(2, 1, 3)
|
||||
with self.assertRaisesRegex(ValueError, 'missing.*top'):
|
||||
encoder.embed_prefix(images=images, state=state, task=['a', 'b'])
|
||||
images['top'] = torch.rand(2, 1, 3, 2, 2)
|
||||
with self.assertRaisesRegex(ValueError, 'task batch'):
|
||||
encoder.embed_prefix(images=images, state=state, task=['only one'])
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user