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:
@@ -422,3 +422,45 @@ class TrainVLATransformerOptimizerTest(unittest.TestCase):
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
||||
class TrainVLASmolVLAOptimizerTest(unittest.TestCase):
|
||||
def test_build_training_optimizer_excludes_frozen_vlm_parameters_and_keeps_state_proj_and_head(self):
|
||||
module = TrainVLATransformerOptimizerTest()._load_train_vla_module()
|
||||
|
||||
class _Head(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(2, 2)
|
||||
|
||||
def get_optim_groups(self, weight_decay):
|
||||
return [{'params': list(self.parameters()), 'weight_decay': weight_decay}]
|
||||
|
||||
class _ConditionEncoder(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.vlm = nn.Linear(2, 2)
|
||||
for param in self.vlm.parameters():
|
||||
param.requires_grad = False
|
||||
self.state_proj = nn.Linear(2, 2)
|
||||
|
||||
class _Agent(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.noise_pred_net = _Head()
|
||||
self.condition_encoder = _ConditionEncoder()
|
||||
|
||||
agent = _Agent()
|
||||
with mock.patch.object(module, 'AdamW', RecordingAdamW):
|
||||
optimizer = module.build_training_optimizer(agent, lr=1e-4, weight_decay=0.01)
|
||||
|
||||
names_by_param_id = {id(param): name for name, param in agent.named_parameters()}
|
||||
optimizer_names = {
|
||||
names_by_param_id[id(param)]
|
||||
for group in optimizer.param_groups
|
||||
for param in group['params']
|
||||
}
|
||||
self.assertIn('condition_encoder.state_proj.weight', optimizer_names)
|
||||
self.assertIn('condition_encoder.state_proj.bias', optimizer_names)
|
||||
self.assertIn('noise_pred_net.proj.weight', optimizer_names)
|
||||
self.assertNotIn('condition_encoder.vlm.weight', optimizer_names)
|
||||
self.assertNotIn('condition_encoder.vlm.bias', optimizer_names)
|
||||
|
||||
Reference in New Issue
Block a user