
    ^j$                    ~    d Z ddlmZ ddlZddlZddlmZ ddlmZm	Z	 ddl
Z
ddlmZmZmZ ddlmZ  G d d	e      Zy)
zAExponential Moving Average callback compatible with ``ModelEma``.    )annotationsN)deepcopy)AnyOptional)CallbackLightningModuleTrainer)AveragedModelc                       e Zd ZdZ	 	 	 	 d	 	 	 	 	 	 	 	 	 d fdZ	 	 	 	 	 	 	 	 ddZddZ	 	 d	 	 	 	 	 ddZddZ	 	 	 	 	 	 	 	 	 	 	 	 ddZ	ddZ
dd	Zdd
ZddZddZddZddZ xZS )RFDETREMACallbacka/  Exponential Moving Average with optional tau-based warm-up.

    Drop-in replacement for ``rfdetr.util.utils.ModelEma`` implemented as a plain Lightning callback around
    :class:`torch.optim.swa_utils.AveragedModel`. The ``_avg_fn`` reproduces the exact same formula as ``ModelEma``
    (1-indexed ``updates`` counter, optional ``tau`` warm-up).

    Args:
        decay: Base EMA decay factor. Corresponds to ``TrainConfig.ema_decay``.
        tau: Warm-up time constant (in optimizer steps). When > 0 the
            effective decay ramps from 0 towards *decay* following ``decay * (1 - exp(-updates / tau))``. Corresponds to
            ``TrainConfig.ema_tau``.
        use_buffers: Whether buffers are averaged in addition to parameters.
        update_interval_steps: Update EMA every N optimizer steps.
    c                    t         |           || _        || _        || _        t        dt        |            | _        d | _        d| _	        d| _
        d | _        d | _        y )N   r   )super__init___decay_tau_use_buffersmaxint_update_interval_steps_average_model_latest_update_step_latest_update_epoch_swapped_state_dict_pending_average_state_dict)selfdecaytauuse_buffersupdate_interval_steps	__class__s        h/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/rfdetr/training/callbacks/ema.pyr   zRFDETREMACallback.__init__$   sc     		'&)!S1F-G&H#7;#$ $&!FJ EI(    c                    |dz   }| j                   dkD  r4| j                  dt        j                  | | j                   z        z
  z  }n| j                  }||z  |d|z
  z  z   S )aN  Compute the EMA update for a single parameter tensor.

        Matches the ``ModelEma`` formula where ``updates`` is 1-indexed: PTL's ``num_averaged`` starts at 0 (incremented
        *after* calling ``avg_fn``), so ``updates = num_averaged + 1`` reproduces the same sequence of effective decay
        values.

        Args:
            averaged_param: Current EMA parameter value.
            model_param: Corresponding live model parameter value.
            num_averaged: Number of models averaged so far (0-indexed).

        Returns:
            Updated EMA parameter tensor.
        r   r   g      ?)r   r   mathexp)r   averaged_parammodel_paramnum_averagedupdateseffective_decays         r#   _avg_fnzRFDETREMACallback._avg_fn7   sc    ( "99q="kkQ7(TYY:N1O-OPO"kkO/+AV2WWWr$   c                   |dk7  ryt        ||j                  | j                  | j                        | _        | j                  j                          | j                  -| j                  j                  | j                         d| _        yt        |d      rt        |d      }t        |t              r| j                  j                  j                  j                  |d      }|j                  s|j                  rIt!        j"                  dt%        |j                         dt%        |j                         d	t&        d
       t)        |d       yy)a  Initialise the averaged model at fit start.

        Args:
            trainer: The Lightning Trainer instance.
            pl_module: The ``RFDETRModelModule`` being trained.
            stage: Current trainer stage (``"fit"``, ``"validate"``, ...).
        fitN)modeldevicer    avg_fn_pending_legacy_ema_stateFstrictz?Legacy EMA checkpoint loaded with non-exact key match; missing=z unexpected=.   )
stacklevel)r
   r1   r   r-   r   evalr   load_state_dicthasattrgetattr
isinstancedictmoduler0   missing_keysunexpected_keyswarningswarnlenUserWarningdelattr)r   trainer	pl_modulestagelegacy_ema_stateincompatibles         r#   setupzRFDETREMACallback.setupR   s2    E>+##))<<	
 	  "++7//0P0PQ/3D,Y ;<&y2MN*D1#2299??OOP`inOo,,0L0LMM##&|'@'@#A"B C&&),*F*F&G%HK $#$ I:; =r$   c                    |duxs |duS )a  Return ``True`` after every optimizer step and every epoch end.

        The base ``WeightAveraging`` only updates on steps. This override also triggers an update at epoch boundaries,
        matching RF-DETR's existing EMA behaviour.

        Args:
            step_idx: Index of the last optimizer step, or ``None``.
            epoch_idx: Index of the last epoch, or ``None``.

        Returns:
            Whether the averaged model should be updated.
        N )r   step_idx	epoch_idxs      r#   should_updatezRFDETREMACallback.should_updatey   s    " t#<y'<<r$   c                &   | j                   y| j                  Tt        |j                               | _        |j	                  | j                   j
                  j                         d       y|j	                  | j                  d       d| _        y)z2Swap live model weights with averaged EMA weights.NTr4   )r   r   r   
state_dictr:   r?   )r   rH   s     r#   _swap_modelszRFDETREMACallback._swap_models   s    &##+'/	0D0D0F'GD$%%d&9&9&@&@&K&K&MVZ%[!!$":":4!H#' r$   c                ,   | j                   y|j                  dz
  }|j                  | j                  k  ry|j                  | _        |j                  | j                  z  dk(  }|r/| j	                  |      r| j                   j                  |       yyy)z!Update EMA after optimizer steps.Nr   r   )rO   )r   global_stepr   r   rQ   update_parameters)r   rG   rH   outputsbatch	batch_idxrO   should_update_steps           r#   on_train_batch_endz$RFDETREMACallback.on_train_batch_end   s     &&&*$":"::#*#6#6 $0043N3NNRSS$"4"4h"4"G11)< #Hr$   c                    | j                   y|j                  | j                  kD  rJ| j                  |j                        r-| j                   j	                  |       |j                  | _        yyy)z*Optionally update EMA at epoch boundaries.N)rP   )r   current_epochr   rQ   rW   r   rG   rH   s      r#   on_train_epoch_endz$RFDETREMACallback.on_train_epoch_end   sh    &  4#<#<<ASAS^e^s^sASAt11)<(/(=(=D% Bu<r$   c                &    | j                  |       y)z*Evaluate tests using averaged EMA weights.NrT   r_   s      r#   on_test_epoch_startz%RFDETREMACallback.on_test_epoch_start       )$r$   c                &    | j                  |       y)z+Restore live weights after test evaluation.Nrb   r_   s      r#   on_test_epoch_endz#RFDETREMACallback.on_test_epoch_end   rd   r$   c                    | j                   5|j                  | j                   j                  j                         d       d| _        y)z6Leave the module in EMA state after training finishes.NTr4   )r   r:   r?   rS   r   r_   s      r#   on_train_endzRFDETREMACallback.on_train_end   s?    *%%d&9&9&@&@&K&K&MVZ%[#' r$   c                    | j                   | j                  d}| j                  | j                  j                         |d<   |S )z(Return callback state for checkpointing.)latest_update_steplatest_update_epochaverage_model_state_dict)r   r   r   rS   )r   states     r#   rS   zRFDETREMACallback.state_dict   sJ     #'":":#'#<#<
 *040C0C0N0N0PE,-r$   c                    |j                  dd      | _        |j                  dd      | _        |j                  d      | _        y)z(Restore callback state from checkpoints.rj   r   rk   r   rl   N)getr   r   r   )r   rS   s     r#   r:   z!RFDETREMACallback.load_state_dict   s<    #->>2F#J $.NN3H"$M!+5>>:T+U(r$   c                @   | j                    t        | j                   j                  d      sy| j                   j                  j                  j	                         j                         D ci c]$  \  }}||j                         j                         & c}}S c c}}w )z;Expose EMA model weights for external checkpoint callbacks.Nr0   )r   r;   r?   r0   rS   itemsdetachclone)r   kvs      r#   get_ema_model_state_dictz*RFDETREMACallback.get_ema_model_state_dict   sw    &gd6I6I6P6PRY.Z262E2E2L2L2R2R2]2]2_2e2e2gh$!Q188:##%%hhhs   -)B)g-?d   Tr   )
r   floatr   r   r    boolr!   r   returnNone)r(   torch.Tensorr)   r|   r*   r   rz   r|   )rG   r	   rH   r   rI   strrz   r{   )NN)rO   Optional[int]rP   r~   rz   ry   )rH   r   rz   r{   )rG   r	   rH   r   rX   r   rY   r   rZ   r   rz   r{   )rG   r	   rH   r   rz   r{   )rz   dict[str, Any])rS   r   rz   r{   )rz   z!Optional[dict[str, torch.Tensor]])__name__
__module____qualname____doc__r   r-   rL   rQ   rT   r\   r`   rc   rf   rh   rS   r:   rv   __classcell__)r"   s   @r#   r   r      s   "  %&JJ J 	J
  #J 
J&X$X "X 	X
 
X6%<R #'#'== != 
	=&	(== #= 	=
 = = 
=(>%%(Vir$   r   )r   
__future__r   r&   rB   copyr   typingr   r   torchpytorch_lightningr   r   r	   torch.optim.swa_utilsr
   r   rN   r$   r#   <module>r      s6    H "       @ @ /Ai Air$   