robometer/rewardeval_ui
4
1from typing import Dict, Any2 3import numpy as np4 5from rfm.data.dataset_types import PreferenceSample, Trajectory6from rfm.data.samplers.base import RFMBaseSampler7 8 9class BaseQualityPreferenceSampler(RFMBaseSampler):10 """Base class for quality preference samplers.11 12 Subclasses should implement `_generate_all_sample_indices` to define how13 trajectories are paired. This base class provides the common `_generate_sample_from_indices`14 method that loads and processes the trajectories.15 """16 17 def _generate_sample_from_indices(self, sample_idx_info: Dict[str, Any]) -> PreferenceSample:18 """Generate a single sample from stored indices."""19 chosen_idx = sample_idx_info["chosen_traj_idx"]20 rejected_idx = sample_idx_info["rejected_traj_idx"]21 22 # Get the trajectories23 chosen_traj = self.dataset[chosen_idx]24 rejected_traj = self.dataset[rejected_idx]25 26 chosen_metadata = {27 "quality_label": chosen_traj["quality_label"],28 "data_source": chosen_traj["data_source"],29 "task": chosen_traj["task"],30 "id": chosen_traj["id"],31 "video_path": chosen_traj["frames"],32 }33 # Add partial_success if available34 if chosen_traj.get("partial_success") is not None:35 chosen_metadata["partial_success"] = chosen_traj.get("partial_success")36 37 chosen_trajectory = self._get_traj_from_data(38 traj=chosen_traj,39 metadata=chosen_metadata,40 )41 42 rejected_metadata = {43 "quality_label": rejected_traj["quality_label"],44 "data_source": rejected_traj["data_source"],45 "task": rejected_traj["task"],46 "id": rejected_traj["id"],47 "video_path": rejected_traj["frames"],48 }49 # Add partial_success if available50 if rejected_traj.get("partial_success") is not None:51 rejected_metadata["partial_success"] = rejected_traj.get("partial_success")52 53 rejected_trajectory = self._get_traj_from_data(54 traj=rejected_traj,55 metadata=rejected_metadata,56 )57 58 data_gen_strategy = getattr(self, "data_gen_strategy", "quality_preference")59 60 # Create preference sample61 sample = PreferenceSample(62 chosen_trajectory=chosen_trajectory,63 rejected_trajectory=rejected_trajectory,64 data_gen_strategy=data_gen_strategy,65 )66 67 return sample68 69 def __len__(self):70 return len(self.sample_indices)71 72 def __getitem__(self, idx):73 return self._generate_sample_from_indices(self.sample_indices[idx])74 