Download scripts/test_widowx_memory.py from DanTim05/DynaTTT: direct link, hf CLI and curl.
- Browser
- Download file 7.55 kB
-
https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/test_widowx_memory.py
- Command line
-
hf download hf://spaces/DanTim05/DynaTTT/scripts/test_widowx_memory.py
-
curl -L -o test_widowx_memory.py https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/test_widowx_memory.py
7.55 kB
| """CPU checks for an external-bank-only ablation (not all-Zeva disabled).""" | |
| import copy | |
| import json | |
| from pathlib import Path | |
| import tempfile | |
| from types import SimpleNamespace | |
| import unittest | |
| from unittest.mock import patch | |
| import numpy as np | |
| import torch | |
| class MemoryAblationTests(unittest.TestCase): | |
| def setUp(self): | |
| self._grad_enabled=torch.is_grad_enabled() | |
| torch.set_grad_enabled(True) | |
| def tearDown(self): | |
| # The upstream inference package disables grad globally on import. | |
| torch.set_grad_enabled(self._grad_enabled) | |
| def test_bank_off_loads_only_cte_and_keeps_policy_weights(self): | |
| from scripts.widowx_policy_server import WidowXService | |
| from cosmos_framework.model.zeva import CausalTransitionEncoder, CausalTransitionEncoderConfig | |
| cfg=CausalTransitionEncoderConfig(action_dim=7,hidden_dim=32,num_layers=1,num_heads=4) | |
| cte=CausalTransitionEncoder(cfg) | |
| model=torch.nn.Linear(2,2) | |
| model.net=SimpleNamespace(behavior_pbd=SimpleNamespace(cfg=SimpleNamespace(global_dim=256))) | |
| before={k:v.clone() for k,v in model.state_dict().items()} | |
| service=object.__new__(WidowXService) | |
| service._memory_mode='bank_off'; service.model=model; service.cfg=SimpleNamespace(action_dim=7) | |
| args=SimpleNamespace(cte_checkpoint=Path('/cte.pt'),bit_mode='normal', | |
| task_context_bank=Path('/must_not_load_bank.pt'), | |
| static_task_context_checkpoint=Path('/must_not_load_head.pt')) | |
| with patch('scripts.widowx_policy_server.torch.load',return_value={ | |
| 'model_config':cfg.to_dict(),'model':cte.state_dict()}) as load: | |
| service._init_zeva(args) | |
| load.assert_called_once_with(args.cte_checkpoint,map_location='cpu',weights_only=False) | |
| self.assertTrue(service._zeva_enabled) | |
| self.assertFalse(service._static_task_context_enabled) | |
| self.assertIsNotNone(service._cte) | |
| with patch('cosmos_framework.scripts.action_policy_server_robocasa365_zeva.retrieve_static_task_context', | |
| side_effect=AssertionError('bank retrieval must not run')): | |
| for prompt in ('put carrot on plate','open microwave'): | |
| value=service._task_context_from_initial_observation(np.zeros((2,2,3),np.uint8),prompt) | |
| self.assertEqual(value.shape,(1,256)) | |
| self.assertEqual(torch.count_nonzero(value).item(),0) | |
| for name,value in model.state_dict().items(): | |
| torch.testing.assert_close(value,before[name],rtol=0,atol=0) | |
| def test_memory_on_uses_unchanged_release_initialization(self): | |
| from scripts.widowx_policy_server import WidowXService, RobolabPolicyService | |
| service=object.__new__(WidowXService);service._memory_mode='on' | |
| args=object() | |
| with patch.object(RobolabPolicyService,'_init_zeva') as original: | |
| service._init_zeva(args) | |
| original.assert_called_once_with(args) | |
| def test_frozen_context_changes_trainable_gradients(self): | |
| from cosmos_framework.model.zeva.policy_injection import PolicyInjectionPrior, PolicyInjectionConfig, gaussian_prior_nll | |
| torch.manual_seed(7) | |
| a=PolicyInjectionPrior(PolicyInjectionConfig(action_dim=7,horizon=32,hidden_dim=32,num_heads=4)) | |
| b=copy.deepcopy(a) | |
| context=torch.randn(2,256,requires_grad=False) | |
| phase=torch.randn(2,128);effects=torch.randn(2,4,128);valid=torch.ones(2,4,dtype=torch.bool) | |
| target=torch.randn(2,32,7) | |
| for model,g in ((a,context),(b,torch.zeros_like(context))): | |
| mean,std=model(g,phase,effects,valid) | |
| gaussian_prior_nll(target,mean,std).backward() | |
| self.assertIsNone(context.grad) | |
| self.assertGreater(a.global_to_anchors.weight.grad.abs().sum().item(),0) | |
| self.assertEqual(b.global_to_anchors.weight.grad.abs().sum().item(),0) | |
| class AblationControllerTests(unittest.TestCase): | |
| def report(checkpoint, mode, horizon): | |
| return dict(status='complete', memory_mode=mode, checkpoint=str(checkpoint), | |
| max_steps=150, inference_num_steps=30, exec_horizon=horizon, | |
| guidance=3.0, seed_start=7, success_rate=0.0, | |
| episodes=[dict(seed=seed, success=False, progress={}) for seed in (7, 8)]) | |
| def test_horizon_16_is_forwarded_to_both_modes_and_recorded(self): | |
| from scripts import run_widowx_memory_ablation as controller | |
| with tempfile.TemporaryDirectory() as directory: | |
| root = Path(directory) | |
| checkpoint = root/'step_000290000.pt'; checkpoint.write_bytes(b'test checkpoint') | |
| output = root/'comparison' | |
| modes = [] | |
| def simulate_watcher(command, env, check): | |
| self.assertTrue(check) | |
| self.assertEqual(command[command.index('--exec-horizon')+1], '16') | |
| self.assertEqual(command[command.index('--num-steps')+1], '30') | |
| self.assertEqual(command[command.index('--max-steps')+1], '150') | |
| mode = command[command.index('--memory-mode')+1]; modes.append(mode) | |
| run = Path(env['WIDOWX_RUN_DIR'])/'eval'; run.mkdir(parents=True) | |
| (run/f'{checkpoint.stem}.json').write_text(json.dumps(self.report(checkpoint, mode, 16))) | |
| traces = run/'traces'/checkpoint.stem; traces.mkdir(parents=True) | |
| for seed in (7, 8): | |
| (traces/f'episode_{seed:02d}.json').write_text(json.dumps([{'command': [0.0]*6+[1.0]}]*150)) | |
| with patch.object(controller.sys, 'argv', ['ablation', '--checkpoint', str(checkpoint), | |
| '--output', str(output), '--exec-horizon', '16', | |
| '--episodes', '2', '--seed-start', '7']), \ | |
| patch.object(controller.subprocess, 'run', side_effect=simulate_watcher): | |
| controller.main() | |
| result = json.loads((output/'comparison.json').read_text()) | |
| self.assertEqual(modes, ['bank_off', 'on']) | |
| self.assertEqual(result['protocol']['exec_horizon'], 16) | |
| self.assertEqual(result['results']['bank_off']['control_steps'], 300) | |
| self.assertEqual(result['status'], 'complete') | |
| def test_four_step_report_cannot_be_used_as_sixteen_step_result(self): | |
| from scripts import run_widowx_memory_ablation as controller | |
| with tempfile.TemporaryDirectory() as directory: | |
| root = Path(directory) | |
| checkpoint = root/'step_000290000.pt'; checkpoint.write_bytes(b'test checkpoint') | |
| output = root/'comparison' | |
| run = output/'bank_off/eval'; run.mkdir(parents=True) | |
| (run/f'{checkpoint.stem}.json').write_text(json.dumps(self.report(checkpoint, 'bank_off', 4))) | |
| with patch.object(controller.sys, 'argv', ['ablation', '--checkpoint', str(checkpoint), | |
| '--output', str(output), '--exec-horizon', '16', | |
| '--episodes', '2', '--seed-start', '7']), \ | |
| patch.object(controller.subprocess, 'run') as launch: | |
| with self.assertRaisesRegex(ValueError, 'exec_horizon'): | |
| controller.main() | |
| launch.assert_not_called() | |
| self.assertFalse((output/'comparison.json').exists()) | |
| if __name__=='__main__': unittest.main() | |