anhtld commited on
Commit
1095311
·
verified ·
1 Parent(s): 6bef905

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,