feat(train): align native SmolVLA recipe
This commit is contained in:
@@ -39,6 +39,17 @@ class FakeLoader:
|
||||
return iter(())
|
||||
|
||||
|
||||
class FakeTqdm:
|
||||
def __init__(self, iterable, **_kwargs):
|
||||
self.iterable = iterable
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self.iterable)
|
||||
|
||||
def set_postfix(self, *_args, **_kwargs):
|
||||
return None
|
||||
|
||||
|
||||
class FakeScheduler:
|
||||
def state_dict(self):
|
||||
return {}
|
||||
@@ -46,13 +57,18 @@ class FakeScheduler:
|
||||
def load_state_dict(self, state_dict):
|
||||
return None
|
||||
|
||||
def step(self):
|
||||
return None
|
||||
|
||||
|
||||
class RecordingAdamW:
|
||||
created = []
|
||||
|
||||
def __init__(self, params, lr, weight_decay):
|
||||
def __init__(self, params, lr, weight_decay, betas=(0.9, 0.999), eps=1e-8):
|
||||
self.lr = lr
|
||||
self.weight_decay = weight_decay
|
||||
self.betas = betas
|
||||
self.eps = eps
|
||||
self.param_groups = self._normalize_param_groups(params, lr, weight_decay)
|
||||
RecordingAdamW.created.append(self)
|
||||
|
||||
@@ -79,6 +95,12 @@ class RecordingAdamW:
|
||||
def load_state_dict(self, state_dict):
|
||||
return None
|
||||
|
||||
def zero_grad(self):
|
||||
return None
|
||||
|
||||
def step(self):
|
||||
return None
|
||||
|
||||
|
||||
class RecordingTransformerHead(nn.Module):
|
||||
def __init__(self):
|
||||
@@ -324,7 +346,7 @@ class TrainVLATransformerOptimizerTest(unittest.TestCase):
|
||||
mock.patch.object(module, 'get_lr_schedule_with_warmup', return_value=FakeScheduler()), \
|
||||
mock.patch.object(module, 'AdamW', RecordingAdamW), \
|
||||
mock.patch.object(module.torch, 'save', return_value=None), \
|
||||
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: iterable):
|
||||
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: FakeTqdm(iterable, **kwargs)):
|
||||
module.main(cfg)
|
||||
finally:
|
||||
os.chdir(previous_cwd)
|
||||
@@ -404,7 +426,7 @@ class TrainVLATransformerOptimizerTest(unittest.TestCase):
|
||||
mock.patch.object(module, 'get_lr_schedule_with_warmup', return_value=FakeScheduler()), \
|
||||
mock.patch.object(module, 'AdamW', RecordingAdamW), \
|
||||
mock.patch.object(module.torch, 'save', return_value=None), \
|
||||
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: iterable):
|
||||
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: FakeTqdm(iterable, **kwargs)):
|
||||
module.main(cfg)
|
||||
finally:
|
||||
os.chdir(previous_cwd)
|
||||
@@ -464,3 +486,117 @@ class TrainVLASmolVLAOptimizerTest(unittest.TestCase):
|
||||
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)
|
||||
|
||||
def test_smolvla_native_training_preset_overrides_optimizer_scheduler_and_grad_clip(self):
|
||||
module = TrainVLATransformerOptimizerTest()._load_train_vla_module()
|
||||
|
||||
class _NativeAgent(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.model = nn.Linear(2, 2)
|
||||
|
||||
def to(self, device):
|
||||
return self
|
||||
|
||||
def get_normalization_stats(self):
|
||||
return {}
|
||||
|
||||
agent = _NativeAgent()
|
||||
cfg = TrainVLATransformerOptimizerTest()._make_cfg()
|
||||
cfg.agent = AttrDict(_target_='roboimi.vla.agent_smolvla_native.SmolVLANativeAgent')
|
||||
cfg.train.lr = 9e-4
|
||||
cfg.train.max_steps = 1
|
||||
cfg.train.weight_decay = 0.123
|
||||
cfg.train.grad_clip = 1.0
|
||||
cfg.train.warmup_steps = 7
|
||||
cfg.train.scheduler_type = 'constant'
|
||||
cfg.train.min_lr = 1e-7
|
||||
|
||||
scheduler_calls = []
|
||||
clip_calls = []
|
||||
|
||||
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_scheduler(*args, **kwargs):
|
||||
scheduler_calls.append(kwargs)
|
||||
return FakeScheduler()
|
||||
|
||||
def fake_clip(parameters, max_norm):
|
||||
clip_calls.append(float(max_norm))
|
||||
return torch.tensor(0.0)
|
||||
|
||||
class OneBatchLoader:
|
||||
def __len__(self):
|
||||
return 1
|
||||
|
||||
def __iter__(self):
|
||||
batch = {
|
||||
'observation.front': torch.zeros(1, 1, 1, 2, 2),
|
||||
'observation.state': torch.zeros(1, 1, 2),
|
||||
'action': torch.zeros(1, 1, 2),
|
||||
}
|
||||
return iter([batch])
|
||||
|
||||
def fake_compute_loss(_batch):
|
||||
return agent.model.weight.sum() * 0.0 + torch.tensor(1.0, requires_grad=True)
|
||||
|
||||
agent.compute_loss = fake_compute_loss
|
||||
|
||||
with tempfile.TemporaryDirectory() as tempdir:
|
||||
previous_cwd = os.getcwd()
|
||||
try:
|
||||
os.chdir(tempdir)
|
||||
with mock.patch.object(module, 'instantiate', side_effect=fake_instantiate), \
|
||||
mock.patch.object(module, 'DataLoader', side_effect=lambda *args, **kwargs: OneBatchLoader()), \
|
||||
mock.patch.object(module, 'get_lr_schedule_with_warmup', side_effect=fake_scheduler), \
|
||||
mock.patch.object(module, 'AdamW', RecordingAdamW), \
|
||||
mock.patch.object(module.torch.nn.utils, 'clip_grad_norm_', side_effect=fake_clip), \
|
||||
mock.patch.object(module.torch, 'save', return_value=None), \
|
||||
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: FakeTqdm(iterable, **kwargs)):
|
||||
module.main(cfg)
|
||||
finally:
|
||||
os.chdir(previous_cwd)
|
||||
|
||||
optimizer = RecordingAdamW.created[-1]
|
||||
self.assertEqual(optimizer.lr, 1e-4)
|
||||
self.assertEqual(optimizer.weight_decay, 1e-10)
|
||||
self.assertEqual(optimizer.betas, (0.9, 0.95))
|
||||
self.assertEqual(optimizer.eps, 1e-8)
|
||||
self.assertEqual(scheduler_calls[-1], {
|
||||
'warmup_steps': 1000,
|
||||
'max_steps': 1,
|
||||
'scheduler_type': 'cosine',
|
||||
'min_lr': 2.5e-6,
|
||||
})
|
||||
self.assertEqual(clip_calls, [10.0])
|
||||
|
||||
def test_cosine_scheduler_spans_requested_training_steps_then_clamps(self):
|
||||
module = TrainVLATransformerOptimizerTest()._load_train_vla_module()
|
||||
param = nn.Parameter(torch.tensor(1.0))
|
||||
optimizer = torch.optim.SGD([param], lr=1e-4)
|
||||
base_lr = optimizer.param_groups[0]['lr']
|
||||
|
||||
scheduler = module.get_lr_schedule_with_warmup(
|
||||
optimizer,
|
||||
warmup_steps=1000,
|
||||
max_steps=150000,
|
||||
scheduler_type='cosine',
|
||||
min_lr=2.5e-6,
|
||||
)
|
||||
|
||||
lr_lambda = scheduler.lr_lambdas[0]
|
||||
observed = {}
|
||||
for step in (0, 1, 1000, 30000, 40000, 150000, 160000):
|
||||
observed[step] = base_lr * lr_lambda(step)
|
||||
|
||||
self.assertGreater(observed[1], observed[0])
|
||||
self.assertLess(observed[1000], 1e-4)
|
||||
self.assertGreater(observed[30000], observed[40000])
|
||||
self.assertGreater(observed[40000], observed[150000])
|
||||
self.assertAlmostEqual(observed[150000], 2.5e-6, places=12)
|
||||
self.assertAlmostEqual(observed[160000], 2.5e-6, places=12)
|
||||
|
||||
Reference in New Issue
Block a user