
    ^jI              	           d Z ddlmZmZmZmZ ddlmZ ddlm	Z	 ddl
mZ  e       Zddededed	efd
Zddeded	efdZdedej$                  d	eeeef      fdZy)zFunctions to get params dict.    )AnyDictListcastN)Joiner)
get_loggernamelr_decay_rate
num_layersreturnc                 *   |dz   }| j                  d      rEd| v sd| v rd}n:d| v r6d| vr2t        | | j                  d      d j                  d	      d
         dz   }t        j                  dj                  | ||dz   |z
  z               ||dz   |z
  z  S )zCalculate lr decay rate for different ViT blocks.

    Args:
        name: parameter name.
        lr_decay_rate: base lr decay rate.
        num_layers: number of ViT blocks.

    Returns:
        lr decay rate for the given parameter.
       backbonez
.pos_embedz.patch_embedr   z.blocks.z
.residual.N.   zname: {}, lr_decay: {})
startswithintfindsplitloggerdebugformat)r	   r
   r   layer_ids       g/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/rfdetr/training/param_groups.pyget_vit_lr_decay_rater      s     A~Hz"4>T#9H4L$<4		* 5 78>>sCAFG!KH
LL)00}VWZbIb7cdeZ!^h677    weight_decay_ratec                 |    d| v sd| v sd| v sd| v sd| v rd}t         j                  dj                  | |             |S )zCalculate weight decay rate for different ViT parameters.

    Args:
        name: parameter name.
        weight_decay_rate: base weight decay rate.

    Returns:
        weight decay rate for the given parameter.
    gamma	pos_embedrel_posbiasnormg        zname: {}, weight_decay rate: {})r   r   r   )r	   r   s     r   get_vit_weight_decay_rater$   *   sP     	4[D0i46GVW[^agkoao
LL299$@QRSr   argsmodel_without_ddpc                    t        |j                  t              sJ t        t        |j                  d         }|j                  | d      }|j                         D cg c]  \  }}|	 }}}d}|j                         D 	cg c]  \  }}	||v s|	j                  s|	 }
}}	|
D cg c]  }|| j                  | j                  z  d  }}|j                         D 	cg c]  \  }}	||vr||vr|	j                  r|	 }}}	|D cg c]  }|| j                  d }}||z   |z   }|S c c}}w c c}	}w c c}w c c}	}w c c}w )Nr   z
backbone.0)prefixztransformer.decoder)paramslr)
isinstancer   r   r   r   get_named_param_lr_pairsitemsnamed_parametersrequires_gradr*   lr_component_decay)r%   r&   r   backbone_named_param_lr_pairs_
param_dictbackbone_param_lr_pairsdecoder_keynpdecoder_paramsparamdecoder_param_lr_pairsother_paramsother_param_dictsfinal_param_dictss                   r   get_param_dictr>   :   s`   '00&999C*33A67H$,$E$EdS_$E$`!?\?b?b?demazee'K$5$F$F$HqDAqK[\L\abapapaqNqftu]bdgg@W@W6WXuu &668Aq22{!7KPQP_P_ 	
L 
 HTTeE9TT),CCF\\! f ru
 Us*    D*D0D0D0)#D6! D;E)      ?   )r?   )__doc__typingr   r   r   r   torch.nnnnrfdetr.models.backboner   rfdetr.utilities.loggerr   r   strfloatr   r   r$   Moduler>    r   r   <module>rK      s    $ ( (  ) .	8 8E 8S 8Z_ 8*C E E    tDcN?S r   