zhicao commited on
Commit
d8686ad
·
verified ·
1 Parent(s): ec36079

Scheduled Commit

Browse files
Files changed (1) hide show
  1. scripts/eval/trex_ablation_eval.py +22 -4
scripts/eval/trex_ablation_eval.py CHANGED
@@ -47,6 +47,7 @@ from groot.vla.model.trex_track_force.force import (
47
  )
48
  from groot.vla.model.trex_track_force.runtime import TrexRuntimeStatistics
49
  from groot.vla.model.trex_track_force.track import TRACK_HORIZON
 
50
  from groot.vla.model.trex_track_force.vla import TrexTrackForceVLA
51
 
52
  ACTION_RATE_HZ = 20.0
@@ -496,10 +497,27 @@ def main() -> None:
496
  for anchor_index, anchor in enumerate(openloop_anchor_times):
497
  raw = builder.build(anchor, blocks=1, prompt=prompt,
498
  with_deform=use_deform, history_only_video=True)
499
- sample = transform_eval(dict(raw))
500
- gt = np.asarray(sample["action"], dtype=np.float32)[..., :62]
501
- gt = torch.as_tensor(gt.reshape(-1, 62)[:ACTION_HORIZON])
502
- batch = to_batch(sample, collator, device, torch.bfloat16)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
503
  if use_deform and "deform_current" in batch:
504
  batch["deform_current"] = batch["deform_current"][:, :1, :1]
505
  modes = [("cascade", True)] if use_force else []
 
47
  )
48
  from groot.vla.model.trex_track_force.runtime import TrexRuntimeStatistics
49
  from groot.vla.model.trex_track_force.track import TRACK_HORIZON
50
+ from groot.vla.model.n1_5.sim_policy import unsqueeze_dict_values
51
  from groot.vla.model.trex_track_force.vla import TrexTrackForceVLA
52
 
53
  ACTION_RATE_HZ = 20.0
 
497
  for anchor_index, anchor in enumerate(openloop_anchor_times):
498
  raw = builder.build(anchor, blocks=1, prompt=prompt,
499
  with_deform=use_deform, history_only_video=True)
500
+ delta = np.asarray(raw.pop("action.eef62"), dtype=np.float32)
501
+ scale = builder.stats.action_q99 - builder.stats.action_q01
502
+ gt = np.clip(
503
+ 2.0 * (delta - builder.stats.action_q01)
504
+ / np.where(scale == 0, 1.0, scale)
505
+ - 1.0,
506
+ -1.0,
507
+ 1.0,
508
+ )
509
+ gt = np.where(scale == 0, delta, gt)
510
+ gt = torch.as_tensor(gt[:ACTION_HORIZON])
511
+ collated = transform_eval(unsqueeze_dict_values(dict(raw)))
512
+ batch = {}
513
+ for key, value in collated.items():
514
+ if torch.is_tensor(value):
515
+ value = (
516
+ value.to(device=device, dtype=torch.bfloat16)
517
+ if value.is_floating_point()
518
+ else value.to(device=device)
519
+ )
520
+ batch[key] = value
521
  if use_deform and "deform_current" in batch:
522
  batch["deform_current"] = batch["deform_current"][:, :1, :1]
523
  modes = [("cascade", True)] if use_force else []