CoolFace
Apppublic

robometer/rewardeval_ui

sourceHugging Faceupdated 7mo agoView on Hugging Face
4likes
base_pref.py74 linesDownload Raw Back to eval
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