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