
    ^j                     B   d Z ddlZddlZddlZddlmZ ddlZddlZddl	m
Z ddlmc mZ ddlmZ ddlmZ ddlmZ ddlmZ ddlmZmZmZ dd	lmZmZmZm Z  dd
l!m"Z" ddl#m$Z$m%Z%m&Z&m'Z'm(Z( ddl)m*Z*m+Z+m,Z, ddl-m.Z.  e.       Z/de0de0fdZ1dedefdZ2 G d de      Z3y)u>   COCOEvalCallback — torchmetrics-based mAP and F1 evaluation.    N)Any)Callback)MeanAveragePrecision)get_coco_api_from_dataset)sweep_confidence_thresholds)DEFAULT_KEYPOINT_MAX_DETSMetricKeypointOKSOKSKey)build_matching_datadistributed_merge_matching_datainit_matching_accumulatormerge_matching_data)box_cxcywh_to_xyxy)_IS_RICH_AVAILABLE_get_rich_console_has_progress_bar_render_overall_merged_render_summary_tables)
all_gatherget_world_sizeis_dist_avail_and_initialized)
get_loggerwarning_emittedreturnc                 4    | ryt         j                  d       y)zWarn once when metric table rendering is skipped because Rich is unavailable.

    Args:
        warning_emitted: Whether this warning has already been emitted.

    Returns:
        Always ``True``; caller assigns back to suppress future warnings.
    TzXRich is not installed; skipping metric table rendering. Install `rich` to enable tables.)loggerwarning)r   s    n/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/rfdetr/training/callbacks/coco_eval.py_warn_missing_rich_oncer   /   s     
NNmn    ema_cbc                 B    | yt        | dd      }|yt        |d|      S )ur  Return the inner ``nn.Module`` wrapped by an EMA callback.

    ``RFDETREMACallback._average_model`` is a private attribute holding a ``torch.optim.swa_utils.AveragedModel``
    (which exposes the actual module on ``.module``).  This helper centralises the access so that consumers degrade
    gracefully when the EMA model has not yet been initialised — preferable to reaching through two layers of
    private attributes at every call site.

    Args:
        ema_cb: EMA callback instance (or ``None``).

    Returns:
        The inner module wrapped by ``AveragedModel``, or ``None`` when no EMA model is available.
    N_average_modelmodule)getattr)r!   averageds     r   _get_ema_inner_moduler'   >   s3     ~v/6H8Xx00r    c                       e Zd ZdZedddddfdededed	ed
ee   dz  dedz  ddf fdZ	de
de
deddfdZde
de
deddfdZde
de
ddfdZde
de
ddfdZde
de
ddfdZde
de
de
de
deddfdZde
de
ddfdZde
de
deee
f   de
deddfdZde
de
ddfdZ	 d@de
de
deee
f   de
dededdfdZde
de
ddfdZdd de
de
d!ed"e
dz  ddf
d#Zd!eddfd$Zde
de
fd%Zde
d"e
deee
f   fd&Zde
de
ddfd'Zde
defd(Zd)Zd"e
ddfd*Ze d"e
defd+       Z!de
d!ede"dz  fd,Z#d!eddfd-Z$de
deee
f   d!eddfd.Z%dd/d0d!ede
de
d1edz  d2eddfd3Z&d4eee
f   d5ed!ede
d6eeef   d7eeeeef   f   deeee
f      fd8Z'de
d!ed9eeef   d:eeee
f      ddf
d;Z(d<eeee)jT                  f      deeee)jT                  f      fd=Z+d>eeee)jT                  f      deeee)jT                  f      fd?Z, xZ-S )ACOCOEvalCallbacka  Validation callback that computes mAP (via torchmetrics) and macro-F1.

    Accumulates predictions and targets across validation batches, then at epoch end computes:

    - ``val/mAP_50_95``, ``val/mAP_50``, ``val/mAP_75``, ``val/mAR`` using
      ``torchmetrics.detection.MeanAveragePrecision``.
    - Per-class ``val/AP/<name>`` when class names are available.
    - ``val/F1``, ``val/precision``, ``val/recall`` from a confidence-threshold
      sweep over compact per-class matching data (DDP-safe).

    For segmentation models (``segmentation=True``) additional metrics ``val/segm_mAP_50_95`` and ``val/segm_mAP_50``
    are logged.

    Args:
        max_dets: Maximum detections per image passed to
            ``MeanAveragePrecision``. Defaults to :data:`~rfdetr.evaluation.keypoint_oks.DEFAULT_KEYPOINT_MAX_DETS`.
        segmentation: When ``True``, evaluate both bbox and segm IoU using
            ``backend="faster_coco_eval"``. Defaults to ``False``.
        eval_interval: Run validation metrics every N epochs. Test metrics are
            always computed when ``trainer.test()`` is called.
        log_per_class_metrics: When ``False``, skip per-class AP logging/table.
    F   TNmax_detssegmentationeval_intervallog_per_class_metricskeypoint_oks_sigmasin_notebookr   c                     t         |           || _        || _        t	        dt        |            | _        t        |      | _        g | _	        i | _
        t               | _        t               | _        d| _        d| _        d | _        d| _        || _        d| _        i | _        || _        d| _        |7t/        j0                  t2              5  ddlm}  |       d u| _        d d d        y || _        y # 1 sw Y   y xY w)Nr*   Fr   )get_ipython)super__init__	_max_dets_segmentationmaxint_eval_intervalbool_log_per_class_metrics_class_names_cat_id_to_namer   	_f1_local_f1_train_local_ema_has_updates_missing_rich_warning_emitted_output_widget_keypoint_mode_use_segm_metrics_train_segm_skip_warned_keypoint_oks_metrics_keypoint_oks_sigmas_in_notebook
contextlibsuppressImportErrorIPythonr2   )	selfr+   r,   r-   r.   r/   r0   r2   	__class__s	           r   r4   zCOCOEvalCallback.__init__l   s     	!)!!S%78&*+@&A#')/14M4O:S:U ',38*#'$)'3-2$CE"$7!"'$$[1 >/$/M$=!> >
 !,D> >s   C44C=trainer	pl_modulestagec                    t        |dd      }|t        |dd      nd}|du | _        | j                  xr | j                   | _        | j                  rddgnd}t	        ddd	| j
                  gd
      }d|d<   t        dd|i|| _        t        dd|i|| _        | j                  j                  j                         D 	ch c]  \  }}	t        |	t              s| }
}}	t        | j                        }|
|k7  rKt        d| j                  j                   j"                   dt%        |
|z
         dt%        ||
z
         d      d| _        yc c}	}w )zInstantiate ``MeanAveragePrecision`` after DDP device placement.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
            stage: One of ``"fit"``, ``"validate"``, ``"test"``, ``"predict"``.
        model_configNuse_grouppose_keypointsFTbboxsegmr*   
   )class_metricsmax_detection_thresholdssync_on_computefaster_coco_evalbackendiou_typezZCOCOEvalCallback._MAP_STATE_ATTRS is out of sync with the installed torchmetrics (version z"). Missing from _MAP_STATE_ATTRS: z. Stale in _MAP_STATE_ATTRS: z. Re-run: python -c "from torchmetrics.detection import MeanAveragePrecision; m = MeanAveragePrecision(); print(sorted(k for k, v in m._defaults.items() if isinstance(v, list)))" and update COCOEvalCallback._MAP_STATE_ATTRS to match. )r%   rC   r6   rD   dictr5   r   
map_metricmap_metric_train	_defaultsitems
isinstancelistset_MAP_STATE_ATTRSRuntimeErrorrN   
__module__sortedmap_metric_ema)rM   rO   rP   rQ   rS   rT   r]   kwargskv	installeddeclareds               r   setupzCOCOEvalCallback.setup   s    y.$? HTG_GL";UCej 	  6=!%!3!3!OD<O<O8O,0,B,B(!%&'T^^%< "
"
 /y.KKFK 4 Qh Q& Q $(??#<#<#B#B#D\41a
STVZH[Q\	\t,,- !__66AAB C339)h:N3O2P Q//5h6J/K.L MJJ	 	 $(! ]s   EEc                     d| _         y)zRelease the notebook output widget when the trainer exits.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
            stage: One of ``"fit"``, ``"validate"``, ``"test"``, ``"predict"``.
        N)rB   )rM   rO   rP   rQ   s       r   teardownzCOCOEvalCallback.teardown   s     #r    c                 \   |j                   }|yt        |d      r|j                  xs g | _        dD ]  }t	        ||d      }|t	        |dd      }|#t        |d      s0t        |d      rE|j
                  j                         D ci c]  \  }}||j                  |   d    c}}| _         y|j                  j                         D 	
ci c]  \  }	}
|	|
d    c}
}	| _         y t        | j                        D ci c]  \  }}||
 c}}| _        yc c}}w c c}
}	w c c}}w )u  Pull class names from the DataModule once the datasets are set up.

        Builds a ``category_id → name`` mapping from the COCO annotation metadata so that per-class AP is logged under
        the class name regardless of whether the dataset uses sequential or non-sequential category IDs.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
        Nclass_names)_dataset_train_dataset_valcococats	label2catname)

datamodulehasattrru   r<   r%   rz   rc   ry   r=   	enumerate)rM   rO   rP   dmattrdatasetrx   labelcat_idrm   rn   ir{   s                r   on_fit_startzCOCOEvalCallback.on_fit_start   s.    :2}% " 4"D6 	Db$-G7FD1DGD&$94-
 OSnnNbNbNd,=JUFtyy088,D(  FJYY__EV+WTQAqyL+WD(!	$ 8AARAR7STGAt4T,
 ,X  Us   DD"D(c                     | j                   j                          t               | _        | j	                  d       | j	                  d       | j                  ||       y)zPrepare the EMA metric on every rank before validation (keeps DDP collectives symmetric).

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
        valval_emaNr`   resetr   r>   _reset_keypoint_split_prepare_ema_metricrM   rO   rP   s      r   on_validation_epoch_startz*COCOEvalCallback.on_validation_epoch_start   sJ     	24""5)""9-  )4r    c                     | j                   j                          t               | _        | j	                  d       | j                  ||       y)a  Reset ``_ema_has_updates`` before test to prevent stale validation state from triggering EMA compute.

        ``on_test_batch_end`` never sets ``_ema_has_updates = True``, so EMA compute is always skipped during
        test (test metrics already reflect the EMA model via checkpoint loading in
        :class:`~rfdetr.training.callbacks.best_model.BestModelCallback`).  Without this hook a stale ``True`` value
        left by a preceding validation epoch would make ``_should_compute_ema`` return ``True``, causing an
        empty-state EMA compute pass that logs sentinel ``-1`` values.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
        testNr   r   s      r   on_test_epoch_startz$COCOEvalCallback.on_test_epoch_start  s<     	24""6*  )4r    outputsbatch	batch_idxc                    t        t        |dd      dd      duryt        |t              rd|vsd|vry| j                  |d         }| j	                  |d         }| j
                  r2|r0d|d	   vr)| j                  st        j                  d
       d| _        y| j                  j                  ||       | j
                  rdnd}t        ||d|      }	t        | j                  |	       | j                  ||d       y)aa  Accumulate train predictions for optional train-split mAP logging.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
            outputs: Return value of ``training_step``.
            batch: The device-transferred batch (unused here).
            batch_idx: Batch index within the training epoch.
        train_configNcompute_train_metricsFTresultstargetsmasksr   zTrain-split segmentation mAP skipped: pred_masks is a sparse dict during training (sparse_forward).  Only val/test segm mAP is available.rV   rU         ?iou_thresholdr]   trainsplit)r%   rd   r_   _convert_preds_convert_targetsrD   rE   r   infora   updater   r   r?   _update_keypoint_oks_metric)
rM   rO   rP   r   r   r   predsr   r]   batch_matchings
             r   on_train_batch_endz#COCOEvalCallback.on_train_batch_end  s   " 79nd;=TV[\dhh'4(IW,D	Y`H`/3/B/B79CU/V''	(:; !!euQx0G//N 04,$$UG4!336,UG3YabD00.A(('(Ir    c                 0   t        t        |dd      dd      dur;| j                  j                          t               | _        | j                  d       y| j                  dkD  rt        t        |dd	            dz   }t        |d
d      }t        |t              xr |d	kD  xr ||k\  }|| j                  z  d	k7  r=|s;| j                  j                          t               | _        | j                  d       y| j                  ||d| j                         y)zCompute optional train-split mAP at the end of the training epoch.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
        r   Nr   FTr   r*   current_epochr   
max_epochsmetric)
r%   ra   r   r   r?   r   r9   r8   rd   _compute_and_logrM   rO   rP   r   r   is_last_epochs         r   on_train_epoch_endz#COCOEvalCallback.on_train_epoch_end>  s    79nd;=TV[\dhh!!'')#<#>D &&w/"! DEIM ,=J&z37jJNj}`jOjMt222a7%%++-'@'B$**73gy'$BWBWXr    c                 n   | j                  |d         }| j                  |d         }| j                  j                  ||       | j                  rdnd}t        ||d|      }	t        | j                  |	       | j                  ||d       | j                  |      }
t        |
      }|
|| j                  |\  }}t        j                  |d   D cg c]  }|d
   	 c}      j                  |j                        }|j                   }t        j"                         5  |j%                           ||      }|j'                  ||      }d	d	d	       | j                        }| j                  j                  ||       | j                  |||d   dd       d| _        y	y	y	y	c c}w # 1 sw Y   `xY w)a  Accumulate predictions and matching data for one validation batch.

        Expects ``outputs`` to be the dict returned by ``RFDETRModelModule.validation_step``: ``{"results": list[dict],
        "targets": list[dict]}``.

        When an EMA callback is present the EMA model is run on the same batch in a separate ``torch.no_grad()`` forward
        pass so that base and EMA metrics are computed from independent predictions.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
            outputs: Return value of ``validation_step``.
            batch: The device-transferred batch ``(samples, targets)``.
            batch_idx: Batch index within the validation epoch.
        r   r   rV   rU   r   r   r   r   N	orig_size)r   r   r   T)r   r   r`   r   rD   r   r   r>   r   _get_ema_callbackr'   rk   torchstacktodevicemodelno_gradevalpostprocessr@   )rM   rO   rP   r   r   r   r   r   r]   r   r!   	ema_innersamples_t
orig_sizesema_underlyingema_outputsema_results	ema_predss                       r   on_validation_batch_endz(COCOEvalCallback.on_validation_batch_endU  s   . 04/B/B79CU/V''	(:;ug.!336,UG3YabDNNN;(('(G ''0)&1	)"7D<O<O<[JGQgi>P%Qan%QRUUV_VfVfgJ&__N M##%,W5'33KLM ++K8I&&y':,,'GI4FG - 
 %)D! =\"7%QM Ms   F&!+F++F4c                    | j                   dkD  rt        t        |dd            dz   }t        |dd      }t        |t              xr |dkD  xr ||k\  }|| j                   z  dk7  rt|sr| j                  j                          | j                  | j                  j                          t               | _        | j                  d       | j                  d       y| j                  ||d       y)zCompute and log mAP and F1 metrics at the end of the validation epoch.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
        r*   r   r   r   Nr   r   )r9   r8   r%   rd   r`   r   rk   r   r>   r   r   r   s         r   on_validation_epoch_endz(COCOEvalCallback.on_validation_epoch_end  s     "! DEIM ,=J&z37jJNj}`jOjMt222a7%%'&&2''--/!:!<**51**95gy%8r    dataloader_idxc                    | j                  |d         }| j                  |d         }| j                  j                  ||       | j                  rdnd}	t        ||d|	      }
t        | j                  |
       | j                  ||d       y	)
a  Accumulate predictions and matching data for one test batch.

        Mirrors :meth:`on_validation_batch_end` for the test evaluation loop triggered by ``trainer.test()`` at the end
        of training.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
            outputs: Return value of ``test_step``.
            batch: Raw batch (unused here).
            batch_idx: Batch index within the test epoch.
            dataloader_idx: Index of the test dataloader (unused here).
        r   r   rV   rU   r   r   r   r   N)	r   r   r`   r   rD   r   r   r>   r   )rM   rO   rP   r   r   r   r   r   r   r]   r   s              r   on_test_batch_endz"COCOEvalCallback.on_test_batch_end  s    , 04/B/B79CU/V''	(:;ug.!336,UG3YabDNNN;(('(Hr    c                 *    | j                  ||d       y)a   Compute and log mAP and F1 under ``test/`` prefix at end of test epoch.

        Mirrors :meth:`on_validation_epoch_end` for the test evaluation loop.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
        r   N)r   r   s      r   on_test_epoch_endz"COCOEvalCallback.on_test_epoch_end  s     	gy&9r    r   r   r   c                b   || j                   n|}|dk(  r| j                  n| j                  }| j                  |      sI|j	                          | j                  |       | j                  |       t        j                  d|       y| j                  |       | j                  ||      }| j                  rdnd}| d| j                   }dt        || d         d	t        || d
         dt        || d         d| j                   t        ||         i}	|j                  | d|| d   dddd       |j                  | d|| d
   dddd       |j                  | d|| d   ddd       |j                  | d||   ddd       || d   j                         j!                         |j"                  | d<   || d
   j                         j!                         |j"                  | d<   || d   j                         j!                         |j"                  | d<   ||   j                         j!                         |j"                  | d<   | j%                  |      }
|
r| j                  | j&                         | j                  || j&                        }|j                  | d|| d   dddd       |j                  | d|| d
   ddd       |j                  | d||   ddd       || d   j                         j!                         |j"                  | d<   || d
   j                         j!                         |j"                  | d<   ||   j                         j!                         |j"                  | d<   | j                  r|j                  | d|d   ddd       |j                  | d|d   ddd       |d   j                         j!                         |j"                  | d<   |d   j                         j!                         |j"                  | d<   | j&                  j	                          n&| j&                  | j&                  j	                          | j                  rt        |d         |	d<   t        |d         |	d<   |j                  | d|d   ddd       |j                  | d |d   ddd       |d   j                         j!                         |j"                  | d<   |d   j                         j!                         |j"                  | d <   t)        |      }i }|rt+        |j-                               }|D cg c]  }||   	 }}t/        |      D cg c]  \  }}||   d!   d"kD  s| }}}t1        |t3        j4                  d"d#d$      |      }t7        |d% &      }t        |d'         |	d(<   t        |d)         |	d*<   t        |d+         |	d,<   |j                  | d-t        |d'         dddd       |j                  | d.t        |d)         ddd       |j                  | d/t        |d+         ddd       t9        j:                  t        |d'               |j"                  | d-<   t9        j:                  t        |d)               |j"                  | d.<   t9        j:                  t        |d+               |j"                  | d/<   t/        |      D ];  \  }}t        |d0   |         t        |d1   |         t        |d2   |         d3||<   = nd4|	d(<   d4|	d*<   d4|	d,<   |j                  | d-d4dddd       |j                  | d.d4ddd       |j                  | d/d4ddd       t9        j:                  d4      |j"                  | d-<   t9        j:                  d4      |j"                  | d.<   t9        j:                  d4      |j"                  | d/<   d5|v r|d5   j<                  d"k(  rt?        |      }|d5   jA                  d"      |d5<   tC        |      D ]O  }tE        ||   t8        jF                        s!||   j<                  d"k(  s4d6|v s9||   jA                  d"      ||<   Q | d| j                   d7}i }||v r5d5|v r1tI        |d5   ||         D ]  \  }}t        |      |tK        |      <    | jM                  ||||||8      }| jO                  |||	|       | jQ                  |||       |d9k(  r|
r| jQ                  d:||d9d;<       n|d9k(  r| j                  d:       |j	                          | j                  |       yc c}w c c}}w )=u  Shared epoch-end logic for validation and test evaluation loops.

        Computes mAP (via ``self.map_metric``), runs the F1 confidence-threshold sweep, logs all scalar metrics via
        ``pl_module.log``, prints two summary tables to the terminal, and resets internal accumulators.  When
        ``self.map_metric_ema`` is set, EMA variants of all metrics (including ``ema_segm_mAP_50_95`` and
        ``ema_segm_mAP_50`` for segmentation models) are logged under the same ``split/`` namespace.

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule.
            split: Metric namespace — ``"val"`` or ``"test"``.
            metric: Optional split-specific mAP accumulator. Defaults to the validation/test accumulator.
        Nr   zHSkipping %s COCO metric compute because no predictions were accumulated.bbox_ mar_z	mAP 50:95mapzmAP 50map_50zmAP 75map_75zmAR @z
/mAP_50_95TFprog_barr   on_stepon_epochz/mAP_50z/mAP_75)r   r   r   z/mARz/ema_mAP_50_95z/ema_mAP_50z/ema_mARz/ema_segm_mAP_50_95segm_mapz/ema_segm_mAP_50segm_map_50zsegm mAP 50:95zsegm mAP 50z/segm_mAP_50_95z/segm_mAP_50total_gtr   r*   e   c                     | d   S )Nmacro_f1r^   )xs    r   <lambda>z3COCOEvalCallback._compute_and_log.<locals>.<lambda>F  s
    : r    )keyr   F1macro_precision	Precisionmacro_recallRecallz/F1z
/precisionz/recallper_class_f1per_class_precper_class_recf1	precisionrecallg        classes	per_class
_per_class)metricspfxr   rP   	ar_by_cid	f1_by_cidr   r   ema_	log_splitmetric_prefix))r`   r?   r>   _metric_has_updatesr   _reset_f1_localr   r   debug _merge_metric_state_across_ranks_compute_map_metricrD   r5   floatlogdetachcpucallback_metrics_should_compute_emark   r   rj   keysr~   r   nplinspacer7   r   tensorndimr_   	unsqueezere   rd   Tensorzipr8   _build_per_class_rows_print_metrics_tables_compute_and_log_keypoint_map)rM   rO   rP   r   r   f1_localr   r   mar_keyoverallshould_compute_emaema_metricsmergedr   
sorted_idscidper_class_listr   classes_with_gt
f1_resultsbestrm   	ar_pc_keyr   class_idarr   s                              r   r   z!COCOEvalCallback._compute_and_log  s
    %+N+0G+;4''''/LLN  '&&u-LLcejk
 	--f5**7F; //gREdnn-. w#c{34eGse6N34eGse6N34DNN#$eGG,<&=	%
 	gZ 'SE+"6d\alp 	 	
 	gWw#f~6d\alp 	 	
 	w'C5)@W\gkltngg&6tU]ab :AC59M9T9T9V9Z9Z9\  E7*!566=Vn6M6T6T6V6Z6Z6\  E7'!236=Vn6M6T6T6V6Z6Z6\  E7'!233:73C3J3J3L3P3P3R  E7$0 "55i@11$2E2EF227D<O<OPKMM'(se3K(   MMUG;/uF^1LUYchswMxMMUG8,k'.B4Y^imMnALPSuTW[AYA`A`AbAfAfAhG$$wn%=>>ISEQW.>Y>`>`>b>f>f>hG$$wk%:;;Fw;O;V;V;X;\;\;^G$$wh%78%%g01;z3JSWafqu   g-.M0JSWafqu   KVV`JaJhJhJjJnJnJp((E72E)FGGRS`GaGhGhGjGnGnGp((E72B)CD%%'  , %%'!!(-gj.A(BG$%%*7=+A%BGM"MMUG?3WZ5HQU_dosMtMMUG<0'-2HQU_dosMtBI*BUB\B\B^BbBbBdG$$wo%>??F}?U?\?\?^?b?b?dG$$wl%;< 1:13	.J5?@cfSk@N@/8/DdVQsT^H_bcHcqdOd4^R[[QRTUWZE[]lmJz'>?D!$z"23GDM#(.?)@#AGK  %d>&: ;GHMM'd:&'   MM'$eD1B,C&DT[`ko   MMUG7+U43G-HQU_dosMt6;ll5jIYCZ6[G$$wc]3=B\\%PTUfPgJh=iG$$wj%9::?,,uTR`MaGb:cG$$wg%67#J/ 3^ 4Q 78!&t,<'=a'@!A#D$9!$<="	#  GDM#&GK  #GHMMUG3-tDRWbfMgMMUG:.D%Z^M_MMUG7+SuW[M\6;ll36GG$$wc]3=B\\#=NG$$wj%9::?,,s:KG$$wg%67 GI$6$;$;q$@7mG!(!3!=!=a!@GI'] 9gaj%,,7GAJOOq<PU`deUe!(!5!5a!8GAJ9
 e4/z:	&(	I$8 #GI$6	8J K 5"+09	#h-(5 ..EYR[gp / 
	 	""7E7IF**5)WEE>0..y)WX]ms.te^&&y1U#M Ads   ?f&f+/f+c                 L    |dk(  rt               | _        yt               | _        y)z,Reset the F1 accumulator for a metric split.r   N)r   r?   r>   )rM   r   s     r   r   z COCOEvalCallback._reset_f1_local  s    G#<#>D 68DNr    c                 \    t        |dg       D ]  }t        t        |dd            s|c S  y)z=Return the EMA callback instance, or ``None`` if not present.	callbacksget_ema_model_state_dictN)r%   callable)rM   rO   callbacks      r   r   z"COCOEvalCallback._get_ema_callback  s6    b9 	 H*DdKL	  r    c                 <   t        |      s|j                         S t        t        j                  d      t        j                  d      f}|D cg c]  }||j
                  f }}	 |D ]C  }|j                         t        j                  k  s%|j                  t        j                         E t        j                  t        j                               5  t        j                  t        j                               5  |j                         cddd       cddd       |D ]  \  }}|j                  |        S c c}w # 1 sw Y   nxY wddd       n# 1 sw Y   nxY w|D ]  \  }}|j                  |        y# |D ]  \  }}|j                  |        w xY w)zeCompute a torchmetrics mAP metric while suppressing duplicate terminal summaries under progress bars.r[   zfaster_coco_eval.coreN)r   computer   logging	getLoggerlevelgetEffectiveLevelWARNINGsetLevelrI   redirect_stdoutioStringIOredirect_stderr)rM   rO   r   metric_loggersmetric_loggerprevious_levelsprevious_levels          r   r   z$COCOEvalCallback._compute_map_metric  sr    )>>## '"3"34F"GIZIZ[rIstUcdMM=+>+>?dd	7!/ < 224wF!**7??;< ++BKKM: (J<V<VWYWbWbWd<e (~~'( ( ( 2A 7-~&&~67 e
( ( ( ( ( 2A 7-~&&~67 7-~&&~67sO   D<'&E> AE> (E>E	E	E> E
	E	E> EE> >Fc                 ,   d| _         | j                  |      d| _        y| j                  N| j                  rddgnd}t	        |ddd| j
                  gdd	      j                  |j                        | _        y| j                  j                          y)
a  Ensure ``map_metric_ema`` exists (and is reset) on EVERY rank when EMA is active.

        Driven by the rank-invariant presence of the EMA callback rather than by per-batch state, so any cross-rank
        state merge (via :meth:`_merge_metric_state_across_ranks`) is issued symmetrically across DDP ranks. Previously
        the metric was created lazily in :meth:`on_validation_batch_end`, so a rank with an empty/uneven shard could
        finish without it, skip the merge/compute path, and deadlock validation (#931 / #449).

        Args:
            trainer: The PTL Trainer.
            pl_module: The LightningModule (provides the device for metric placement).
        FNrU   rV   Tr*   rW   r[   )r]   rX   rY   r\   rZ   )	r@   r   rk   rD   r   r5   r   r   r   )rM   rO   rP   ema_iou_types       r   r   z$COCOEvalCallback._prepare_ema_metric  s     !&!!'*2"&D&484J4J 0PVL"6%"*+R)@* %# b!!"  %%'r    c                 F   | j                   duxr | j                  }|rdnd}t               rkt        j                  |gt        |dd            }t        j                  |t        j                  j                         t        |j                               }t        |      S )u  Decide — identically on every rank — whether to run the EMA metric ``compute()``.

        Under DDP, ``_merge_metric_state_across_ranks`` issues cross-rank collectives that every rank must
        participate in, or none may — a rank that skips desynchronises the NCCL collective sequence and deadlocks
        validation (#931 / #449).  Each rank votes ``1`` only when its EMA metric exists and received at least
        one batch update this epoch; a cross-rank ``all_reduce(MIN)`` makes the decision unanimous — a single
        rank voting 0 suppresses EMA compute on all ranks.

        Args:
            pl_module: The LightningModule (provides the device for the reduction).

        Returns:
            ``True`` iff every rank both holds an EMA metric object and received at least one batch update this
            epoch, making ``compute()`` safe to run identically on all ranks; ``False`` otherwise (EMA compute
            skipped uniformly on all ranks).
        Nr*   r   r   r  r   op)rk   r@   r   r   r  r%   dist
all_reduceReduceOpMINr8   itemr:   )rM   rP   has_emavoteflags        r   r  z$COCOEvalCallback._should_compute_ema  s{    " %%T1Kd6K6Kq(*<<wy(E/RSDOODT]]%6%67tyy{#DDzr    )	detection_boxdetection_scoresdetection_labelsdetection_maskgroundtruth_boxgroundtruth_labelsgroundtruth_maskgroundtruth_crowdsgroundtruth_areac                    |t               rt               dk(  ry| j                  D ]  }t        ||d      }||D cg c]7  }t	        j
                  |      r|j                         j                         n|9 }}t        |      }|D cg c]  }|D ]  }|  }	}}t        |||	        t        t        |dd      d      |_        yc c}w c c}}w )u  Merge a metric's accumulated per-rank state onto every rank, replacing torchmetrics' sync.

        torchmetrics' built-in sync (``gather_all_tensors``) varies the number of collectives by each state
        tensor's *local* ndim (scalar → 1 all_gather, vector → 2), so when DDP seg validation leaves a state
        scalar on some ranks and a vector on others the ranks issue different collective counts and deadlock
        (#931 / #449).  Instead we gather each state list once with the repo's pickle-based ``all_gather`` — a
        fixed collective pattern issued identically on every rank regardless of tensor shape — and concatenate.
        With ``sync_on_compute=False`` the metric's own ``compute()`` then runs locally over the merged full-set
        state, yielding the identical global mAP without any shape-dependent collective.

        Args:
            metric: The ``MeanAveragePrecision`` instance whose state should be merged in place.

        Note:
            No-op when ``metric`` is ``None``, when the distributed process group is not
            initialised, or when world size is 1 (single GPU / CPU training).  In these
            cases the metric state is unchanged.
        Nr*   _update_countr   )r   r   rg   r%   r   	is_tensorr   r  r   setattrr7   rL  )
rM   r   r   localrn   	local_cpugathered	rank_listr>  r  s
             r   r   z1COCOEvalCallback._merge_metric_state_across_ranks  s    & >!>!@NDTXYDY)) 		*DFD$/E} QVV1U__Q-?)QFVIV!),H,4KyKdKdKFKFD&)		*  #76?A#FJ WKs   <CCc                     t        | dd      }t        |t              r|dkD  S t        j                  |      r8t        |j                         j                         j                         dkD        S y)zIReturn whether a torchmetrics metric has accumulated at least one update.rL  Nr   T)	r%   rd   r8   r   rM  r:   r   r  r>  )r   update_counts     r   r   z$COCOEvalCallback._metric_has_updates  sa     v=lC(!##??<(++-11388:Q>??r    c                 ^   || j                   v r| j                   |   S t        |dd      }|y|j                  d      }ddddj                  |d      }|D ]T  }t        ||d      }|t	        |      }|!t        || j                  | j                  	      }	|	| j                   |<   |	c S  y)
a}  Return the :class:`~rfdetr.evaluation.keypoint_oks.MetricKeypointOKS` for *split*, creating it if needed.

        The metric is created lazily on first access per split and reused across epochs (state is reset
        at epoch boundaries via :meth:`_reset_keypoint_split`).

        Args:
            trainer: The PTL Trainer (provides access to the datamodule).
            split: One of ``"train"``, ``"val"``, ``"val_ema"``, or ``"test"``.

        Returns:
            A :class:`~rfdetr.evaluation.keypoint_oks.MetricKeypointOKS` bound to the split's COCO
            ground-truth, or ``None`` when no dataset is available.
        r|   N_ema)rv   )rw   )_dataset_test)r   r   r   )rw   rW  rv   )r/   r+   )rF   r%   removesuffixgetr   r	   rG   r5   )
rM   rO   r   r|   source_splitsplit_attrsr   r   coco_apir   s
             r   "_get_or_create_keypoint_oks_metricz3COCOEvalCallback._get_or_create_keypoint_oks_metric!  s     D...--e44WlD9
))&1($&
 #lO
P	 	
   	Dj$5G09H&$($=$=F
 17D&&u-M	 r    c                 `    | j                   j                  |      }||j                          yy)zReset accumulated keypoint predictions for *split*.

        Args:
            split: One of ``"train"``, ``"val"``, ``"val_ema"``, or ``"test"``.
        N)rF   rY  r   )rM   r   r   s      r   r   z&COCOEvalCallback._reset_keypoint_splitL  s.     ++//6LLN r    c                 j   | j                   sy| j                  ||      }|yi }|d   }|d   }t        ||      D ]  \  }}	|	j                  d      }
|
t	        j
                  |
      rt        |
j                               n
t        |
      }d|vri ||<   ]|d   j                         j                         |d   j                         j                         |d   j                         j                         |d   j                         j                         d	||<    |sy|j                  |       y)
a"  Accumulate batch predictions into the keypoint OKS metric.

        Args:
            trainer: The PTL Trainer.
            outputs: Batch output dict with ``"results"`` and ``"targets"`` keys.
            split: Metric split (``"train"``, ``"val"``, ``"val_ema"``, or ``"test"``).
        Nr   r   image_id	keypointsboxesscoreslabels)rb  rc  rd  ra  )rC   r]  r  rY  r   rM  r8   r>  r   r  r   )rM   rO   r   r   r   predictionsr   r   resulttargetimage_id_tensorr`  s               r   r   z,COCOEvalCallback._update_keypoint_oks_metricV  s7    ""88%H>:<)$)$!'73 	NFF$jj4O&6;ooo6Vs?//12\_`o\pH&((*H%//1557 *113779 *113779#K0779==?	%K!	 k"r    r   r   r   r   c          	      Z   | j                   j                  |      }| j                  r|y|j                  rdnd}t	               rkt        j                  |gt        |dd            }t        j                  |t        j                  j                         t        |j                               }|sy||n|}	 |j                         }	t        j                   dft        j"                  dft        j$                  d	ft        j&                  d	fd
}
|
j)                         D ]b  \  }\  }}|	j                  |d      }|dk  r!| d| | }|j+                  |||dd	d       t        j                  |      |j,                  |<   d 	 |j/                          y# |j/                          w xY w)a  Compute and log OKS keypoint AP/AR metrics when keypoint mode is active.

        Args:
            split: Internal metric split (``"val"``, ``"val_ema"``, ``"train"``, ``"test"``).
            pl_module: The LightningModule used to log scalar metrics.
            trainer: The PTL Trainer (provides ``callback_metrics``).
            log_split: Namespace prefix for logged keys. Defaults to *split*.
            metric_prefix: Optional string prepended to each metric name (e.g. ``"ema_"``).
        Nr*   r   r   r  r7  r8  TF)keypoint_map_50_95keypoint_map_50keypoint_map_75keypoint_mARg      /r   )rF   rY  rC   has_updatesr   r   r  r%   r:  r;  r<  r=  r8   r>  r%  r
   MAPMAP_50MAP_75MARrc   r   r  r   )rM   r   rP   rO   r   r   r   has_updates_voterA  statskeypoint_metricsmetric_namestat_keyr   valuelog_keys                   r   r  z.COCOEvalCallback._compute_and_log_keypoint_map{  s   $ ++//6""fn
 !' 2 21(*<<!1 279hX];^_DOODT]]%6%67"499;/&.EI		NN$E'-zz4&8$*MM4#8$*MM5#9!'U 3	  6F5K5K5M H11h		(D119&Kq}EguxV[fjk49LL4G((1H LLNFLLNs   :CF F*r   r   r   r   c                 2   g }| j                   s|S | d}||vsd|vr|S t        |d   ||         D ]  \  }	}
t        |
      }|j                  t	        |	      t        d            }|dk  r||k7  s|dk  rEt	        |	      }| j
                  j                  |t        |            }|j                  | d| |
       |||d}|j                  |j                  |t        d      t        d      t        d      d             |j                  |        |S )a,  Build per-class rows and emit per-class AP metrics.

        Args:
            metrics: Output of ``MeanAveragePrecision.compute()``.
            pfx: Key prefix for bbox metrics when segmentation mode is enabled.
            split: Metric namespace (``"val"`` or ``"test"``).
            pl_module: LightningModule used for metric logging.
            ar_by_cid: Per-class AR keyed by ``category_id``.
            f1_by_cid: Per-class F1/precision/recall keyed by ``category_id``.

        Returns:
            Per-class rows for table rendering.
        map_per_classr   nanr   z/AP/)r{   apr  r   )
r;   r  r   rY  r8   r=   strr   r   append)rM   r   r   r   rP   r   r   r   pc_keyr  r~  ap_far_fidxr{   rows                   r   r  z&COCOEvalCallback._build_per_class_rows  s!   , +-	**5& IW$<	 2GFOD 
	"LHb9D==Xe=DaxTT\TAXh-C''++CS:DMMUG4v.3+/t4"HCJJy}}SuERWLdijodp*qrsS!
	" r    r  r   c                    t        |dd      syt        st        | j                        | _        yt	        |      }t        t        |dd            dz   }t        |dd      }t        |t
              r|dkD  r	d| d	| d
nd| d
}|j                         |z   }	t        |	|| j                        }
| j                  r| j                  St        j                  t              5  ddl}ddlm} |j%                         | _         || j                         ddd       | j                  @| j                  j'                  d       | j                  5  t)        ||	|
|       ddd       yt        j                  t              5  ddlm}  |d       ddd       t)        ||	|
|       yt)        ||	|
|       y# 1 sw Y   xY w# 1 sw Y   yxY w# 1 sw Y   ?xY w)uz  Print two tables to the terminal: overall metrics and per-class metrics.

        The overall table is transposed (metrics as columns, one value row) with true merged group-header cells rendered
        via box-drawing characters: ``mAP`` spans sub-columns 50:95 / 50 / 75, ``mAR`` spans ``@N``, and ``F1 sweep``
        spans F1 / Prec / Recall.  The per-class table uses a standard Rich ``Table`` with columns for AP 50:95, AR, F1,
        Prec, Recall.

        Only runs on the global-zero rank to avoid duplicate output in DDP.

        Args:
            trainer: The PTL Trainer (used to check ``is_global_zero``).
            split: ``"val"`` or ``"test"``.
            overall: Ordered mapping of metric label → scalar value.
            per_class: Per-class dicts with keys ``name``, ``ap``, ``ar``,
                ``f1``, ``precision``, ``recall``; skipped when empty.
        is_global_zeroTNr   r   r*   r   z (Epoch rn  ))display)wait)clear_output)r%   r   r   rA   r   r8   rd   
capitalizer   r5   rH   rB   rI   rJ   rK   
ipywidgetsIPython.displayr  Outputr  r   )rM   rO   r   r  r   consoler   r   	epoch_sfx	title_pfxoverall_renderedwidgetsr  r  s                 r   r  z&COCOEvalCallback._print_metrics_tables  s   . w 0$7!1HIkIk1lD.#G,GG_a@AAEWlD9
 *c*zA~ }oQzl!4M?!, 	
 $$&2	1)WdnnU
 ""*((5 107*1..*:D'D//01 "".##00d0;(( \*7I?OQZ[\ $$[1 (8$'( #7I7GS 	w	3CYO51 1\( (s$    2F5G?G5F>G
Gr   c                     g }|D ]`  }t        |      }d|v r>|d   j                  dk(  r,|d   j                  d   dk(  r|d   j                  d      |d<   |j	                  |       b |S )a  Normalise prediction dicts from ``PostProcess`` for torchmetrics.

        ``PostProcess.forward`` returns masks with shape ``[K, 1, H, W]`` (the extra channel is introduced by
        ``F.interpolate`` which requires 4-D input).  Both ``torchmetrics.MeanAveragePrecision`` and
        ``engine.build_matching_data`` expect ``[K, H, W]``, so squeeze the channel dim when present.

        ``PostProcess.forward`` currently returns ``[K, 1, H, W]`` masks. Keep this callback-local squeeze for metric
        code paths because ``RFDETR.predict`` and other inference-facing callers still consume the 4-D representation
        and apply ``.squeeze(1)`` at their boundary.

        Args:
            preds: Raw per-image prediction dicts from ``PostProcess``.

        Returns:
            Per-image dicts with ``masks`` squeezed to ``[K, H, W]`` when applicable; all other keys are passed through
            unchanged.
        r      r*   )r_   r  shapesqueezer  )rM   r   outpentrys        r   r   zCOCOEvalCallback._convert_preds$  s}    $  	AGE%E'N$7$71$<wAUAUVWAX\]A]!&w!7!7!:gJJu		
 
r    r   c                 4   g }|D ]  }|d   j                         \  }}|d   j                  ||||g      }t        |d         |z  }||d   d}d|v r|d   j                         }	|	j                  dd t        |      t        |      fk7  rft        j                  |	j                         j                  d      t        |      t        |      fd	
      j                  d      j                         }	|	|d<   d|v r|d   |d<   |j                  |        |S )a  Convert targets from normalised CxCyWH to absolute xyxy boxes.

        Also passes ``iscrowd`` and ``masks`` through unchanged.

        Args:
            targets: Per-image target dicts with ``boxes`` in normalised
                CxCyWH format and ``orig_size`` as ``[H, W]``.

        Returns:
            Per-image dicts with ``boxes`` in absolute xyxy, ``labels``, and optionally ``masks`` and ``iscrowd``.
        r   rb  rd  )rb  rd  r   Nr*   nearest)sizemodeiscrowd)tolist
new_tensorr   r:   r  r8   Finterpolater   r	  r  r  )
rM   r   r  r   hwscalerb  r  r   s
             r   r   z!COCOEvalCallback._convert_targets>  s!     	A[>((*DAqgJ))1aA,7E&qz2U:E7<(-TE!|'
) ;;rs#AA'77!KKM33A6"%a&#a&!1!*
 !  "'gA~#$Y<i JJu-	. 
r    )r   ).__name__ri   __qualname____doc__r   r8   r:   re   r   r4   r   r  rq   rs   r   r   r   r   r   r_   r   r   r   r   r   r   r   r   r   r  rg   r   staticmethodr   r	   r]  r   r   r  r  r  r   r
  r   r   __classcell__)rN   s   @r   r)   r)   T   s8   2 2"&*26#'#,#, #, 	#,
  $#, "%[4/#, D[#, 
#,R1(S 1(S 1( 1( 1(f# # #C #D #"UC "UC "UD "UH5 5 5 553 53 54 5$'J'J 'J 	'J
 'J 'J 
'JRY# Y# Y$ Y.9)9) 9) c3h	9)
 9) 9) 
9)v9s 9s 9t 98  II I c3h	I
 I I I 
I@	: 	: 	: 	: bf u$ u$ u$C u$TWZ^T^ u$jn u$n9S 9T 9  73 7 7S#X 7"(3 (3 (4 (8S T B
"Ks "Kt "KH C D  )# )c )N_bfNf )V3 4 ##3 ##c3h ##X[ ##`d ##V !%22 2 	2 :2 2 
2h)c3h) ) 	)
 ) U
#) T#u*--.) 
d38n	)VHPHP HP c5j!	HP
 S#X'HP 
HPTDc5<<.?)@$A d4PSUZUaUaPaKbFc 4$T#u||2C-D(E $$tTWY^YeYeTeOfJg $r    r)   )4r  rI   r-  r&  typingr   numpyr  r   torch.distributeddistributedr:  torch.nn.functionalnn
functionalr  pytorch_lightningr   torchmetrics.detectionr   rfdetr.datasetsr   rfdetr.evaluation.f1_sweepr   rfdetr.evaluation.keypoint_oksr   r	   r
   rfdetr.evaluation.matchingr   r   r   r   rfdetr.utilities.box_opsr   rfdetr.utilities.consoler   r   r   r   r   rfdetr.utilities.distributedr   r   r   rfdetr.utilities.loggerr   r   r:   r   r'   r)   r^   r    r   <module>r     s    E  	         & 7 5 B 
  8  c b .	T d 1# 1# 1,Nx Nr    