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