Scheduled Commit
Browse files
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 |
-
|
| 500 |
-
|
| 501 |
-
gt =
|
| 502 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 []
|