DynaTTT / scripts /audit_widowx_runtime.py
DanTim05's picture
Upload Zeva source and project assets
68eed61 verified
Raw History Blame Contribute Delete
7.1 kB
"""Bounded, isolated diagnostics; never changes or restarts the training job."""
import argparse
import dataclasses
import json
import os
from pathlib import Path
import subprocess
import time
import urllib.request
ROOT = Path(__file__).resolve().parents[1]
COSMOS = Path('/home/gpu4/tianyi/cosmos-framework')
PYTHON = '/home/gpu4/tianyi/cosmos-framework/.venv/bin/python'
SIM = '/home/gpu4/miniconda3/envs/simpler_env/bin/python'
def closed_loop(args):
# Run foreground under this bounded parent; always reap our own server.
with (args.output / 'server.log').open('w') as log:
server = subprocess.Popen([PYTHON, str(ROOT/'scripts/widowx_policy_server.py'),
'--delta', str(args.checkpoint), '--port', str(args.port)],
stdout=log, stderr=subprocess.STDOUT, cwd=COSMOS)
try:
deadline = time.monotonic() + 240
while True:
if server.poll() is not None:
raise RuntimeError(f'server exited {server.returncode}')
try:
with urllib.request.urlopen(f'http://127.0.0.1:{args.port}/healthz', timeout=2) as r:
if r.status == 200:
break
except OSError:
pass
if time.monotonic() >= deadline:
raise TimeoutError('server not ready')
time.sleep(2)
env = dict(os.environ, PYTHONPATH='/opt/benchmarks/SimplerEnv',
MS2_ASSET_DIR='/opt/benchmarks/SimplerEnv/ManiSkill2_real2sim/data',
VK_ICD_FILENAMES='/etc/vulkan/icd.d/nvidia_icd.json')
for convention in ('bridge_tool', 'legacy_tcp'):
destination = args.output / convention
destination.mkdir(exist_ok=True)
with (destination/'episodes.log').open('w') as elog:
subprocess.run([SIM, str(ROOT/'scripts/eval_widowx_simpler.py'),
'--port', str(args.port), '--checkpoint', str(args.checkpoint),
'--output', str(destination/'results.json'), '--episodes', '10',
'--proprio-convention', convention,
'--video-dir', str(destination/'videos'),
'--trace-dir', str(destination/'traces')],
env=env, stdout=elog, stderr=subprocess.STDOUT, check=True, timeout=600,
cwd=COSMOS)
report=json.loads((destination/'results.json').read_text())
print(convention, report['success_rate'], flush=True)
finally:
server.terminate()
try:
server.wait(timeout=20)
except subprocess.TimeoutExpired:
server.kill()
server.wait()
def offline(args):
import numpy as np
import torch
from scripts.widowx_policy_server import WidowXService
from cosmos_framework.scripts.action_policy_server_robocasa365_zeva import RobolabServerArgs
from cosmos_framework.data.generator.action.datasets.widowx_bridge_v3_dataset import WidowXBridgeV3Dataset, normalize_actions
torch.set_num_threads(2)
release=Path('/home/gpu4/tianyi/zeva-release/weights')
os.chdir(COSMOS)
service=WidowXService(RobolabServerArgs(
checkpoint_path=str(release/'stage2'), allow_dcp_checkpoint=True,
experiment='action_policy_robocasa365_atomic5_zeva',
experiment_overrides=['model.config.tokenizer.vae_path='+os.environ['WAN_VAE_PATH'],
'model.config.vlm_config.tokenizer.pretrained_model_name='+os.environ['QWEN_VLM_PATH']],
task_context_bank=release/'stage1/train_memory_effect_v3.pt',
cte_checkpoint=release/'stage1/zeva_cte.pt', static_task_context_checkpoint=release/'stage3/best.pt',
domain_name='bridge_orig_lerobot', image_height=256, image_width=256,
resolution='256', conditioning_fps=5, action_dim=7, use_state=False))
dataset=WidowXBridgeV3Dataset('/opt/zhangchenyu/datasets/bridge_orig_lerobot_smoke',
feature_cache=str(ROOT/'.artifacts/widowx/features'))
indices=[0, dataset.windows.index((306,0)), len(dataset)-1,
next(i for i,(_,s) in enumerate(dataset.windows) if s>=16)]
named=dict(service.model.named_parameters())
metrics=[]
for label, checkpoint in [('step2',ROOT/'.artifacts/widowx/checkpoints/step_000000002.pt'),
('trained',args.checkpoint)]:
payload=torch.load(checkpoint,map_location='cpu',weights_only=False)
with torch.no_grad():
for name,value in payload['model'].items():
p=named[name]; (p.to_local() if hasattr(p,'to_local') else p).copy_(value)
for i in indices:
sample=dataset[i]; ep,start=dataset.windows[i]
frames=dataset._frames[ep]; rows=dataset._rows[ep]
request={'prompt':sample['ai_caption'], 'observation/image':frames[start],
'observation/proprio':sample['proprio'].numpy(),
'cte_boundary_images':frames[:start+1:4],
'cte_transition_actions':normalize_actions(rows['action'][:start]).reshape(-1,4,7),
'inference_seed':0}
for steps,guidance in ((4,3.0),(16,1.0)) if label=='trained' else ((4,3.0),):
service.cfg=dataclasses.replace(service.cfg,num_steps=steps,guidance=guidance)
result=service.infer(request)
prediction=np.asarray(result['action'],np.float32)
target=sample['action'].numpy()
error=np.abs(prediction-target)
item=dict(checkpoint=label,iteration=payload['iteration'],episode=ep,start=start,
steps=steps,guidance=guidance,mae=float(error.mean()),
first4_mae=float(error[:4].mean()),per_dim_mae=error.mean(0).tolist(),
gripper_accuracy=float(np.mean((prediction[:,6]>0)==(target[:,6]>0))),
out_of_range_fraction=float(np.mean(np.abs(prediction)>1)))
metrics.append(item)
np.savez(args.output/f'{label}_ep{ep}_start{start}_steps{steps}.npz',
prediction=prediction,target=target,proprio=sample['proprio'].numpy())
print(json.dumps(item),flush=True)
(args.output/'metrics.json').write_text(json.dumps(metrics,indent=2))
if __name__=='__main__':
parser=argparse.ArgumentParser()
parser.add_argument('--mode',choices=('closed_loop','offline'),required=True)
parser.add_argument('--checkpoint',type=Path,required=True)
parser.add_argument('--output',type=Path,required=True)
parser.add_argument('--port',type=int,default=18766)
args=parser.parse_args(); args.output.mkdir(parents=True,exist_ok=True)
globals()[args.mode](args)