Auto-sync: 2026-06-29 01:34:49
Browse files
dovla_cil/eval/maniskill_policy_rollout.py
CHANGED
|
@@ -83,6 +83,7 @@ def evaluate_maniskill_policy_rollout(
|
|
| 83 |
retrieval_type_min_success: float = 0.0,
|
| 84 |
retrieval_residual_min_source_progress: float = 0.0,
|
| 85 |
retrieval_residual_source_progress_bonus_scale: float = 0.0,
|
|
|
|
| 86 |
retrieval_residual_scale: float = 1.0,
|
| 87 |
retrieval_residual_scales: tuple[float, ...] = (),
|
| 88 |
retrieval_residual_anchor: str = "expert",
|
|
@@ -187,6 +188,8 @@ def evaluate_maniskill_policy_rollout(
|
|
| 187 |
raise ValueError("retrieval_residual_min_source_progress must be in [0, 1]")
|
| 188 |
if retrieval_residual_source_progress_bonus_scale < 0:
|
| 189 |
raise ValueError("retrieval_residual_source_progress_bonus_scale must be non-negative")
|
|
|
|
|
|
|
| 190 |
if retrieval_residual_scale < 0:
|
| 191 |
raise ValueError("retrieval_residual_scale must be non-negative")
|
| 192 |
if any(scale < 0 for scale in retrieval_residual_scales):
|
|
@@ -258,6 +261,9 @@ def evaluate_maniskill_policy_rollout(
|
|
| 258 |
retrieval_residual_source_progress_bonus_scale=(
|
| 259 |
retrieval_residual_source_progress_bonus_scale
|
| 260 |
),
|
|
|
|
|
|
|
|
|
|
| 261 |
retrieval_residual_anchor=retrieval_residual_anchor,
|
| 262 |
retrieval_residual_reduce=retrieval_residual_reduce,
|
| 263 |
)
|
|
@@ -353,6 +359,11 @@ def evaluate_maniskill_policy_rollout(
|
|
| 353 |
if selection_mode == "retrieval_residual"
|
| 354 |
else 0.0
|
| 355 |
),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 356 |
"retrieval_residual_scale": retrieval_residual_scale
|
| 357 |
if selection_mode == "retrieval_residual"
|
| 358 |
else 0.0,
|
|
@@ -537,12 +548,17 @@ def _attach_retrieved_residual_candidates(
|
|
| 537 |
retrieval_type_min_success: float = 0.0,
|
| 538 |
retrieval_residual_min_source_progress: float = 0.0,
|
| 539 |
retrieval_residual_source_progress_bonus_scale: float = 0.0,
|
|
|
|
| 540 |
retrieval_residual_anchor: str = "expert",
|
| 541 |
retrieval_residual_reduce: str = "none",
|
| 542 |
) -> list[_RolloutCase]:
|
| 543 |
if observation_mode != "state":
|
| 544 |
raise ValueError("retrieval_residual currently supports state observations only")
|
| 545 |
heldout = set(heldout_group_ids)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 546 |
type_success_rates = _candidate_type_success_rates(dataset, heldout_group_ids=heldout)
|
| 547 |
bank: dict[
|
| 548 |
str,
|
|
@@ -583,6 +599,7 @@ def _attach_retrieved_residual_candidates(
|
|
| 583 |
continue
|
| 584 |
reward = getattr(record, "reward", None)
|
| 585 |
source_progress = float(getattr(reward, "progress", 0.0))
|
|
|
|
| 586 |
if source_progress < retrieval_residual_min_source_progress:
|
| 587 |
continue
|
| 588 |
residual = np.asarray(_numeric_action_values(record), dtype=np.float32) - anchor_action
|
|
@@ -590,6 +607,7 @@ def _attach_retrieved_residual_candidates(
|
|
| 590 |
candidate_types.append(f"residual_{record.candidate_type}")
|
| 591 |
residual_bonuses.append(
|
| 592 |
float(retrieval_residual_source_progress_bonus_scale) * source_progress
|
|
|
|
| 593 |
)
|
| 594 |
feature = np.asarray(
|
| 595 |
vectorize_toy_observation(records[0].observation_inline or {}, obs_dim=obs_dim),
|
|
@@ -610,9 +628,7 @@ def _attach_retrieved_residual_candidates(
|
|
| 610 |
candidate_action_values=[zero],
|
| 611 |
candidate_types=["policy_residual"],
|
| 612 |
candidate_score_bonuses=(
|
| 613 |
-
[0.0]
|
| 614 |
-
if retrieval_residual_source_progress_bonus_scale > 0
|
| 615 |
-
else None
|
| 616 |
),
|
| 617 |
candidate_source_group_id=None,
|
| 618 |
)
|
|
@@ -664,9 +680,7 @@ def _attach_retrieved_residual_candidates(
|
|
| 664 |
candidate_action_values=residuals,
|
| 665 |
candidate_types=candidate_types,
|
| 666 |
candidate_score_bonuses=(
|
| 667 |
-
residual_bonuses
|
| 668 |
-
if retrieval_residual_source_progress_bonus_scale > 0
|
| 669 |
-
else None
|
| 670 |
),
|
| 671 |
candidate_source_group_id=";".join(source_group_ids),
|
| 672 |
)
|
|
@@ -801,6 +815,15 @@ def _candidate_type_success_rates(
|
|
| 801 |
}
|
| 802 |
|
| 803 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 804 |
def _nearest_retrieval_entries(
|
| 805 |
candidates: list[tuple[Any, np.ndarray, Any, Any]],
|
| 806 |
query: np.ndarray,
|
|
|
|
| 83 |
retrieval_type_min_success: float = 0.0,
|
| 84 |
retrieval_residual_min_source_progress: float = 0.0,
|
| 85 |
retrieval_residual_source_progress_bonus_scale: float = 0.0,
|
| 86 |
+
retrieval_residual_source_score_bonus_scale: float = 0.0,
|
| 87 |
retrieval_residual_scale: float = 1.0,
|
| 88 |
retrieval_residual_scales: tuple[float, ...] = (),
|
| 89 |
retrieval_residual_anchor: str = "expert",
|
|
|
|
| 188 |
raise ValueError("retrieval_residual_min_source_progress must be in [0, 1]")
|
| 189 |
if retrieval_residual_source_progress_bonus_scale < 0:
|
| 190 |
raise ValueError("retrieval_residual_source_progress_bonus_scale must be non-negative")
|
| 191 |
+
if retrieval_residual_source_score_bonus_scale < 0:
|
| 192 |
+
raise ValueError("retrieval_residual_source_score_bonus_scale must be non-negative")
|
| 193 |
if retrieval_residual_scale < 0:
|
| 194 |
raise ValueError("retrieval_residual_scale must be non-negative")
|
| 195 |
if any(scale < 0 for scale in retrieval_residual_scales):
|
|
|
|
| 261 |
retrieval_residual_source_progress_bonus_scale=(
|
| 262 |
retrieval_residual_source_progress_bonus_scale
|
| 263 |
),
|
| 264 |
+
retrieval_residual_source_score_bonus_scale=(
|
| 265 |
+
retrieval_residual_source_score_bonus_scale
|
| 266 |
+
),
|
| 267 |
retrieval_residual_anchor=retrieval_residual_anchor,
|
| 268 |
retrieval_residual_reduce=retrieval_residual_reduce,
|
| 269 |
)
|
|
|
|
| 359 |
if selection_mode == "retrieval_residual"
|
| 360 |
else 0.0
|
| 361 |
),
|
| 362 |
+
"retrieval_residual_source_score_bonus_scale": (
|
| 363 |
+
retrieval_residual_source_score_bonus_scale
|
| 364 |
+
if selection_mode == "retrieval_residual"
|
| 365 |
+
else 0.0
|
| 366 |
+
),
|
| 367 |
"retrieval_residual_scale": retrieval_residual_scale
|
| 368 |
if selection_mode == "retrieval_residual"
|
| 369 |
else 0.0,
|
|
|
|
| 548 |
retrieval_type_min_success: float = 0.0,
|
| 549 |
retrieval_residual_min_source_progress: float = 0.0,
|
| 550 |
retrieval_residual_source_progress_bonus_scale: float = 0.0,
|
| 551 |
+
retrieval_residual_source_score_bonus_scale: float = 0.0,
|
| 552 |
retrieval_residual_anchor: str = "expert",
|
| 553 |
retrieval_residual_reduce: str = "none",
|
| 554 |
) -> list[_RolloutCase]:
|
| 555 |
if observation_mode != "state":
|
| 556 |
raise ValueError("retrieval_residual currently supports state observations only")
|
| 557 |
heldout = set(heldout_group_ids)
|
| 558 |
+
uses_source_bonus = (
|
| 559 |
+
retrieval_residual_source_progress_bonus_scale > 0
|
| 560 |
+
or retrieval_residual_source_score_bonus_scale > 0
|
| 561 |
+
)
|
| 562 |
type_success_rates = _candidate_type_success_rates(dataset, heldout_group_ids=heldout)
|
| 563 |
bank: dict[
|
| 564 |
str,
|
|
|
|
| 599 |
continue
|
| 600 |
reward = getattr(record, "reward", None)
|
| 601 |
source_progress = float(getattr(reward, "progress", 0.0))
|
| 602 |
+
source_score = _source_reward_score(reward, progress=source_progress)
|
| 603 |
if source_progress < retrieval_residual_min_source_progress:
|
| 604 |
continue
|
| 605 |
residual = np.asarray(_numeric_action_values(record), dtype=np.float32) - anchor_action
|
|
|
|
| 607 |
candidate_types.append(f"residual_{record.candidate_type}")
|
| 608 |
residual_bonuses.append(
|
| 609 |
float(retrieval_residual_source_progress_bonus_scale) * source_progress
|
| 610 |
+
+ float(retrieval_residual_source_score_bonus_scale) * source_score
|
| 611 |
)
|
| 612 |
feature = np.asarray(
|
| 613 |
vectorize_toy_observation(records[0].observation_inline or {}, obs_dim=obs_dim),
|
|
|
|
| 628 |
candidate_action_values=[zero],
|
| 629 |
candidate_types=["policy_residual"],
|
| 630 |
candidate_score_bonuses=(
|
| 631 |
+
[0.0] if uses_source_bonus else None
|
|
|
|
|
|
|
| 632 |
),
|
| 633 |
candidate_source_group_id=None,
|
| 634 |
)
|
|
|
|
| 680 |
candidate_action_values=residuals,
|
| 681 |
candidate_types=candidate_types,
|
| 682 |
candidate_score_bonuses=(
|
| 683 |
+
residual_bonuses if uses_source_bonus else None
|
|
|
|
|
|
|
| 684 |
),
|
| 685 |
candidate_source_group_id=";".join(source_group_ids),
|
| 686 |
)
|
|
|
|
| 815 |
}
|
| 816 |
|
| 817 |
|
| 818 |
+
def _source_reward_score(reward: Any, *, progress: float) -> float:
|
| 819 |
+
if reward is None:
|
| 820 |
+
return float(progress)
|
| 821 |
+
score = getattr(reward, "score", None)
|
| 822 |
+
if score is not None:
|
| 823 |
+
return float(score)
|
| 824 |
+
return float(progress) + (1.0 if bool(getattr(reward, "terminal_success", False)) else 0.0)
|
| 825 |
+
|
| 826 |
+
|
| 827 |
def _nearest_retrieval_entries(
|
| 828 |
candidates: list[tuple[Any, np.ndarray, Any, Any]],
|
| 829 |
query: np.ndarray,
|