CoolFace
Apppublic

robometer/rewardeval_ui

sourceHugging Faceupdated 7mo agoView on Hugging Face
4likes
confusion_matrix.cpython-310.pyc100 linesDownload Raw Back to __pycache__
1o

2��ei�3�@sxdZddlZddlZddlmZmZddlmZddlm	Z	m3Z4ddlmZddl
mZddlmZGd	d5�d6e�ZdS)z/7Data generator for confusion matrix analysis.8�N)�Counter�defaultdict)�Tuple)�PreferenceSample�ProgressSample)�RFMBaseSampler)�rank_0_print)�SentenceTransformercs�eZdZdZddef�fdd�
Zdeefdd�Zde	eeeffd	d9�Z10defdd
�Zdedefdd�Z
dd�Zdd�Z�ZS)�ConfusionMatrixSamplera�11    Data generator that creates task-trajectory pairs for confusion matrix analysis.12 13    For each unique task, creates samples with each trajectory to analyze14    how well the model can distinguish between different tasks.15 16    If multiple data sources are present, samples N random trajectories from each data source17    and prioritizes different language instructions by randomizing the pairing order.18    ��n_trajectories_per_sourcecs�t�jdi|��||_td�|_|j��t|j���}t	dt19|��d�|jd�i|_|D]}|j�
|�}t�|�|j|<q/t	dt20|j��d�|jd�|`|��|_t	dt21|j��dt22|j��d	t23|j��d24��dS)
aInitialize confusion matrix sampler.25 26        Args:27            n_trajectories_per_source: Number of trajectories to sample from each data source.28                If None, uses all available trajectories.29            **kwargs: Additional arguments passed to parent class.30        z'sentence-transformers/all-MiniLM-L12-v2z%Precomputing language embeddings for z
 unique tasks��verbosezPrecomputed z language embeddings�31Generated z& confusion matrix sample indices from z trajectories and z tasksN�)�super�__init__rr	Zsentence_model�eval�list�task_indices�keysr�lenr�task_embeddings�encode�torch�tensor�_generate_all_sample_indices�sample_indices�robot_trajectories)�selfr�kwargs�unique_tasks�task�	embedding��	__class__r�B/scr/aliang80/reward_fm/rfm/data/samplers/eval/confusion_matrix.pyrs 323334(�zConfusionMatrixSampler.__init__�returnc35Cs<g}t|j���}tdt|��d|��|jd�|��\}}tdt|��d�|jd�|�|�|��}|j	�36|�t�}|D]-}|j|}|d}	||	d7<|�
dt|��}37|D]}|�|||	|d	|38d39��q\q?|j	�40|�tdt|��d�|jd�td
t|���|jd�tdtt|������|jd�|S)a(Generate all possible task-trajectory pair sample indices.41        42        If multiple data sources exist, samples N random trajectories from each data source.43        Prioritizes different video tasks first, then prioritizes different language instructions44        when creating pairs.45        �Found z unique language tasks: r
zProcessing z+ trajectories for confusion matrix analysisr"��id�frames)�traj_idx�	lang_task�46video_task�47video_pathr*rz task-trajectory pairsz  Video tasks sampled: z  Trajectories per video task: N)rrrrrr�#_sample_trajectories_by_data_source�_print_sampling_stats�copy�
_local_random�shuffler�dataset�get�str�append�dict�sorted�items)rrZunique_lang_tasksZsampled_trajectories�statsZshuffled_lang_tasksZvideo_task_countr,�trajr.�traj_idr-rrr&r=s>�484950��51 z3ConfusionMatrixSampler._generate_all_sample_indicesc	Cs�g}it�id�}tdd��}|jD]}|j|}|�dd�}|�dd�}|||�|�qtdt|��dt|�	����|j52d	�|��D�]0\}}|D]53}	|j�
||	�qMt|�	��}54|j�
|55�td56d�|��D��dd
�|��D�t�d�}|jdur�g}|��D]\}	}
|�|
�t|
�|d|	<|d|	t|
�7<q�td|�dt|��d�|j57d	�n�t|j|d�}g}dd
�|��D�}|58��}d}t|�|k�r(|t|�kr�d}|j�
|�||}	z!t||	�}|�|�|d|	d7<|d|	d7<Wnt�y|�|�|�sY�q(Yq�w|d7}t|�|ks�td|�dt|��d|d�d�|j59d	�tdtt|d������|j60d	�|D]}|j|}|�dt|��}|�dd�|d|<�qQ|�|�||d|<qF||fS)a}Sample N random trajectories from each data source, prioritizing different video tasks.61        62        When sampling N trajectories, first selects one trajectory from each unique video task,63        then repeats in round-robin fashion until N trajectories are sampled.64        65        Returns:66            Tuple of (list of sampled trajectory indices, stats dictionary)67        )�	by_source�by_task�traj_to_taskcSstt�S�N)rrrrrr&�<lambda>�szLConfusionMatrixSampler._sample_trajectories_by_data_source.<locals>.<lambda>�data_source�unknownr"r(z data sources: r
css�|]}t|�VqdSrB�r)�.0�indicesrrr&�	<genexpr>�s�zMConfusionMatrixSampler._sample_trajectories_by_data_source.<locals>.<genexpr>cS�i|]	\}}|t|��qSrrF�rGr"rHrrr&�68<dictcomp>��zNConfusionMatrixSampler._sample_trajectories_by_data_source.<locals>.<dictcomp>)�total_available�tasks_available�
tasks_sampledNrPr@z  Data source 'z
': Using all �
 trajectoriesrNcSrJr)�iterrKrrr&rL�rMrr)z': Sampled z out of z    Tasks sampled: r*rAr?)rrrr5r6r8rrrrrr;r3r4�sum�valuesr�extend�minr2�next�
StopIteration�popr9r:r7)r�sampled_indicesr<Ztrajectories_by_source_and_taskr,r=rDr.Ztasks_to_indicesr"�	all_tasks�source_statsZsampled_from_sourcerHZn_to_sampleZtask_iterators�	task_listZ	round_idxr>rrr&r0{s�	�6970��7172�7374����7576z:ConfusionMatrixSampler._sample_trajectories_by_data_sourcer<c77Cs&|jsdStd|jd�td|jd�t|d���D]\}}td|�d|�d�|jd�qtd	|jd�|d78��D]N\}}td|��|jd�td|d
��|jd�tdt|d���|jd�t|d���D]\}}|d�|d�}td|�d|�d|�d�|jd�qkq;td|jd�dS)z�Print detailed statistics about sampled trajectories.79        80        Args:81            stats: Statistics dictionary from _sample_trajectories_by_data_source82        Nz-83=== Confusion Matrix Sampling Statistics ===r
z%84Overall trajectories per video task:r@z  z: rQz85Per data source breakdown:r?z  Data source: z    Total available: rNz    Tasks available: rOrPrz      �/z trajectories sampledz2==================================================)rrr:r;rr6)rr<r"�countrDr\Z
sampled_countrrr&r1�s&��z,ConfusionMatrixSampler._print_sampling_stats�sample_idx_infocCsz|d}|d}|d}|d}|j|}|j|}|d|||d�}|��}	||	d<||	d<|j|	|d	�}86t|87d88�}|S)z=Generate a single task-trajectory sample from stored indices.r,r-r.r/r*)r*r-r.r/r"�text_embedding)r=�metadata)�89trajectoryN)r5rr2�_get_traj_from_datar)rr`r,r-r.r/Z90video_trajrarbZvideo_traj_with_taskZsample_trajectory�samplerrr&�_generate_sample_from_indicess(9192��93z4ConfusionMatrixSampler._generate_sample_from_indicescCs94t|j�SrB)rr)rrrr&�__len__'s95zConfusionMatrixSampler.__len__cCs|�|j|�SrB)rfr)r�idxrrr&�__getitem__*sz"ConfusionMatrixSampler.__getitem__)r)�__name__�96__module__�__qualname__�__doc__�intrrr9rrr0r1rrfrgri�
__classcell__rrr$r&r97s98!>m r99)rm�randomr�collectionsrr�typingr�rfm.data.dataset_typesrr�rfm.data.samplers.baser�rfm.utils.distributedr�sentence_transformersr	r100rrrr&�<module>s