DynaTTT / scripts /test_widowx_memory.py
DanTim05's picture
Upload Zeva source and project assets
68eed61 verified
Raw History Blame Contribute Delete
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):
@staticmethod
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()