
    ^j                     *   d Z ddlZddlZddlmZmZmZ ddlZddlmZ erddl	m
Z
mZ ddlmZmZ ddlmZ ddlmZ dd	lmZmZmZmZmZmZmZ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% ddl&m'Z' ddl(m)Z)m*Z*  e'       Z+dejX                  de-dejX                  fdZ.dej^                  de-dej^                  fdZ0dejb                  ddfdZ2 G d dejb                        Z3d%dZ4d%dZ5defddd ed!   d"ede3fd#Z6efddd d!d"ede7ee#f   fd$Z8y)&z$LW-DETR model and criterion classes.    N)TYPE_CHECKINGCallableOptional)nn)ModelConfigTrainConfig)MODEL_DEFAULTSModelDefaults)BuilderArgs)build_backbone)SetCriterion	dice_lossdice_loss_jitposition_supervised_losssigmoid_ce_losssigmoid_ce_loss_jitsigmoid_focal_losssigmoid_varifocal_loss)SegmentationHead)build_matcher)MLP)PostProcess)build_transformer)
get_logger)NestedTensornested_tensor_from_tensor_listlinearnum_classesreturnc                 .   | j                   j                  d   }t        t        j                  ||z              }| j                   j                         j                  |d      d| }| j                  ,| j                  j                         j                  |      d| nd}t        j                  | j                  ||du      }t        j                         5  |j                   j                  |       |'|j                  |j                  j                  |       ddd       | j                   j                  |j                   _        | j                  1|j                  %| j                  j                  |j                  _        |S # 1 sw Y   mxY w)u  Return a new :class:`~torch.nn.Linear` resized to *num_classes* outputs.

    Tiles the existing weight rows when *num_classes* is larger than the current output size, or truncates them when
    smaller.  The returned module has ``out_features == num_classes`` so that ``nn.Linear`` metadata stays consistent
    with the actual weight shape — a requirement for correct ONNX export and ``torch.jit.trace`` serialisation.

    Args:
        linear: Source linear layer whose weights are used as the starting point.
        num_classes: Target number of output features.

    Returns:
        A new :class:`~torch.nn.Linear` with ``in_features`` unchanged and ``out_features == num_classes``.
    r      N)bias)weightshapeintmathceildetachrepeatr"   r   Linearin_featurestorchno_gradcopy_requires_grad)r   r   basenum_repeats
new_weightnew_bias
new_linears          _/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/rfdetr/models/lwdetr.py_resize_linearr6   8   s@    ==q!Ddiid 234K%%'..{A>|LJIOI`v{{!!#**;7EfjH6--{QUAUVJ	 ,
+JOO$?OO!!(+, '-mm&A&AJ#{{:??#>(.(A(A
%, ,s   AFF	parameternum_rowsc                    | j                   d   }||k(  r| S |dk(  r1| j                         j                  |g| j                   dd       }nYt        t	        j
                  ||z              gdg| j                         dz
  z  } | j                         j                  | d| }t        j                  |j                         | j                        S )zQReturn a parameter with the first dimension resized by tiling or truncating rows.r   r!   N)r/   )r$   r(   	new_zerosr%   r&   r'   dimr)   r   	Parametercloner/   )r7   r8   current_rowsnew_datarepeatss        r5   _resize_parameter_rowsrA   V   s    ??1%Lxq##%//0PIOOAB<O0PQtyyL!89:[qcY]]_WXEX>Y[,9##%,,g6yA<<(	8O8OPP    modulec                 |   t        | dd      }|t        |      dk(  ry|d   }t        |t        j                        r|j
                  dk  ryt        j                         5  |j                  dd j                          |j                  |j                  dd j                          ddd       y# 1 sw Y   yxY w)zFReset keypoint precision-Cholesky output rows to unit Gaussian values.layersNr            )getattrlen
isinstancer   r*   out_featuresr,   r-   r#   zero_r"   )rC   rE   final_layers      r5   $_reset_keypoint_gaussian_output_rowsrP   d   s    VXt,F~V)*Kk299-1I1IQ1N	 *1Q%%''Qq!'')* * *s   "AB22B;c                       e Zd ZdZ	 	 	 	 	 	 	 	 d$dee   dz  def fdZdeddfdZedee   de	j                  fd	       Zedee   de	j                  fd
       Zdee   fdZedeee	j                  f   dee   dz  fd       Zdee   dz  ddfdZd%dZd Zde	j                  dedede	j                  fdZde	j                  de	j                  fdZd&defdZd Ze	j2                  j4                  	 d&de	j                  de	j                  de	j                  dz  de	j                  dz  fd       Zdeej<                     fdZd e d!eddfd"Z!d# Z" xZ#S )'LWDETRz@This is the Group DETR v3 module that performs object detection.Nnum_keypoints_per_class grouppose_keypoint_dim_downscalec           	      	   t         |           || _        || _        |j                  }t        j                  ||      | _        t        ||dd      | _	        || _
        d}t        j                  ||z  |      | _        t        j                  ||z  |      | _        t
        j                  j                  | j                  j                   j"                  d       || _        || _        || _        |	| _        | j*                  s&| j                  | j                  j,                  _	        nd| j                  j,                  _	        |
| _        || _        |xs g | _        || _        | j0                  r@t7        | j2                        |kD  r(t9        dt7        | j2                         d| d| d      d	| _        | j0                  rt        || j4                  z  || j4                  z  d
d      | _        t
        j                  j                  | j<                  j>                  d   j                   j"                  d       t
        j                  j                  | j<                  j>                  d   j@                  j"                  d       nd| _        | jC                  d| jE                  | j2                               d}tG        jH                  d|z
  |z         }tK        jL                  |      |z  | j                  j@                  _        t
        j                  j                  | j                  j>                  d   j                   j"                  d       t
        j                  j                  | j                  j>                  d   j@                  j"                  d       || _'        | jN                  r t        jP                  tS        |      D cg c]!  }tU        jV                  | j                        # c}      | j                  _,        t        jP                  tS        |      D cg c]!  }tU        jV                  | j                        # c}      | j                  _-        | j0                  rd| j<                  Xt        jP                  tS        |      D cg c]!  }tU        jV                  | j<                        # c}      | j                  _.        d	| _/        yc c}w c c}w c c}w )a  Initializes the model.

        Parameters:
            backbone: torch module of the backbone to be used. See backbone.py
            transformer: torch module of the transformer architecture. See transformer.py
            num_classes: number of object classes
            num_queries: number of object queries, ie detection slot. This is the maximal number of objects
                         Conditional DETR can detect in a single image. For COCO, we recommend 100 queries.
            aux_loss: True if auxiliary decoding losses (loss at each decoder layer) are to be used.
            group_detr: Number of groups to speed detr training. Default is 1.
            lite_refpoint_refine: TODO
        rH      r   Nznum_keypoints_per_class has z) entries but the detection head only has z= classes. Class-logit boosts for keypoint classes with id >= zV would be silently truncated. Increase num_classes or shorten num_keypoints_per_class.F   rF   _kp_active_maskg{Gz?r!   )0super__init__num_queriestransformerd_modelr   r*   class_embedr   
bbox_embedsegmentation_head	Embeddingrefpoint_embed
query_featinit	constant_r#   databackboneaux_loss
group_detrlite_refpoint_refinedecoderbbox_reparamuse_grouppose_keypointsrS   rT   rK   
ValueError_kp_zero_pad_warnedkeypoint_embedrE   r"   register_buffer_create_kp_active_maskr&   logr,   ones	two_stage
ModuleListrangecopydeepcopyenc_out_bbox_embedenc_out_class_embedenc_out_keypoint_embed_export)selfrg   r\   r`   r   r[   rh   ri   ru   rj   rl   rm   rS   rT   
hidden_dim	query_dim
prior_prob
bias_value_	__class__s                      r5   rZ   zLWDETR.__init__w   s   8 	&& ((
99Z=j*a;!2	 ll;+CYO,,{Z'?L
$--44991=  $ %9!((26//D$$/26D$$/( (?$'>'D"$0P-''C0L0L,MP[,[.s43O3O/P.Q R'=(efqer shh  $) ''"%dCCCdCCC	#D GGd1188<CCHH!LGGd1188<AAFFJ"&D.0K0KDLhLh0ij 
hhJ*<==
%*ZZ%<z%I" 	$//004;;@@!D
$//00499>>B #>>24--9>z9JKAt/K3D/ 46==:?
:KLQt//0L4D0 ++0C0C0O:<--AFzARSAT]]4#6#67S;  7  L M Ts   &S&&S&S$r   r   c           	          t        | j                  |      | _        | j                  rQt        j                  | j
                  j                  D cg c]  }t        ||       c}      | j
                  _        yyc c}w )a  Resize the detection classification head to *num_classes* outputs.

        Replaces ``self.class_embed`` (and each ``enc_out_class_embed`` when the model uses two-stage detection) with a
        new :class:`torch.nn.Linear` whose ``out_features`` equals *num_classes*.  When *num_classes* is larger than the
        current head the existing weights are tiled; when smaller they are truncated.  Replacing the module (rather than
        mutating ``.data``) keeps ``nn.Linear.out_features`` consistent with the actual weight shape, which is required
        for correct ONNX export.

        Args:
            num_classes: Target number of output classes (including background).
        N)r6   r^   ru   r   rv   r\   r{   )r~   r   ms      r5   reinitialize_detection_headz"LWDETR.reinitialize_detection_head   s`     *$*:*:KH>>35==9=9I9I9]9]^A;/^4D0 ^s   A:c                    | s&t        j                  ddt         j                        S t        |       }t        j                  t	        |       |t         j                        }t        |       D ]  \  }}d||d|f<    |S )zECreate a compact class-by-keypoint active mask for a keypoint schema.r   dtypeTN)r,   zerosboolmaxrK   	enumerate)rS   max_kp	kp_active	class_idxnum_keypointss        r5   rr   zLWDETR._create_kp_active_mask   sx     ';;q!5::66,-KK$; <fEJJW	(12I(J 	8$I}37Ii-/0	8rB   c                    | s&t        j                  ddt         j                        S t        |       }t        j                  d|z   d|z   t         j                        }t	        |       D ]`  \  }}|dk(  rdt        | d|       z   }||z   }t	        |       D ]1  \  }}|dk(  s||k(  rdt        | d|       z   }	|	|z   }
d||||	|
f<   3 b |S )zGCreate an attention mask that blocks cross-class keypoint interactions.r!   r   r   NT)r,   r   r   sumr   )rS   total_keypointsmaskclass_idx_inum_keypoints_istart_iend_iclass_idx_jnum_keypoints_jstart_jend_js              r5   _create_keypoint_class_maskz"LWDETR._create_keypoint_class_mask   s     ';;q!5::6656{{1.O0C5::V,56M,N 
	:(K!##5l{CDDGo-E09:Q0R :,_"a';++Ec"9,;"GHH/159WU]GEM12:
	: rB   c                     | j                   j                  d      j                         D cg c]  }t        |       c}S c c}w )zJReturn the current keypoint schema inferred from the active-keypoint mask.r!   r;   )rX   r   tolistr%   )r~   r   s     r5   get_num_keypoints_per_classz"LWDETR.get_num_keypoints_per_class  s9    8<8L8L8P8PUV8P8W8^8^8`a}M"aaas   A
state_dictc                     | j                  d      }t        |t        j                        r|j                  dk7  ry|j                  d      j                         D cg c]  }t        |       c}S c c}w )z<Infer the keypoint schema stored in a checkpoint state dict.rX      Nr!   r   )getrL   r,   Tensorndimr   r   r%   )r   active_maskr   s      r5   +get_num_keypoints_per_class_from_checkpointz2LWDETR.get_num_keypoints_per_class_from_checkpoint  s^     !nn%67+u||48H8HA8M8CA8N8U8U8WX}M"XXXs   A3c                 Z   | j                   r|syt        |      }t        |      }|| _        | j	                  |      j                  | j                  j                        | _        t        | j                  d      r|| j                  _        t        | j                  dd      }|t        |d      r||_        t        |dd      }t        |t        j                        rt        ||      |_        t        |d      r|j!                          nGt        |d      r;|j"                  }| j!                  |      j                  |j                        |_        dD ]R  }t        | j                  |d      }t        |dd      }	t        |	t        j                        sBt        |	|      |_        T y)	zMResize schema-dependent GroupPose state to match ``num_keypoints_per_class``.NrS   rk   keypoint_pos_embedr   keypoint_class_mask)keypoint_query_initializerkeypoint_query_initializer_encqueries)rm   listr   rS   rr   torX   devicehasattrr\   rJ   rL   r   r<   rA   r   r   r   r   )
r~   rS   schemar   rk   r   current_maskinitializer_nameinitializerr   s
             r5   reinitialize_keypoint_headz!LWDETR.reinitialize_keypoint_head   sw   ++3J-.f+'-$#::6BEEdFZFZFaFab4##%>?7=D4$**It<w 9:28/!(2F!M,bll;-CDVXg-h*w =>335"78&::.2.N.Nv.V.Y.YZfZmZm.n+ ` 	W!$"2"24DdKKk9d;G'2<<0&<Wo&V#		WrB   c                     | j                   r| j                  yt        | j                         t        | j                  dd      }t        |t        j                        r|D ]  }t        |        yy)a  Reset keypoint Gaussian precision outputs to unit values.

        Keypoint channels 4, 5, and 6 encode the lower-triangular precision
        Cholesky parameters ``log_l11``, ``l21``, and ``log_l22``. Zeroing the
        final prediction rows gives ``L = identity`` at the start of finetuning
        while preserving learned keypoint location, visibility, findability, and
        class-logit channels loaded from the checkpoint.
        Nr|   )rm   rp   rP   rJ   r\   rL   r   rv   )r~   enc_keypoint_embedrp   s      r5   "reset_keypoint_gaussian_parametersz)LWDETR.reset_keypoint_gaussian_parameters@  sm     ++t/B/B/J,T-@-@A$T%5%57OQUV("--8"4 E4^DE 9rB   c                 *   d| _         | j                  | _        | j                  | _        | j	                         D ]W  \  }}t        |d      st        |j                  t              s.t        |d      s;|j                   rH|j                          Y y )NTexportr}   )	r}   forward_forward_originforward_exportnamed_modulesr   rL   r   r   )r~   namer   s      r5   r   zLWDETR.exportR  sr    #||**))+ 	GD!q(#
188X(F7STV_K`ijirir
	rB   keypoints_compactbatch_size_expectednum_queries_expectedc                    | j                   r| j                  s|S |j                         dk7  r|S |j                  \  }}}}||k7  s||k7  rt	        d| d| d| d| d	      t        | j                        }t        | j                        |z  }	t        | j                        }
||	k(  r|S ||
k7  r"t	        d| d|
 d|	 d	| j                   d
	      t        j                  |||	||j                  |j                        }d}t        | j                        D ]>  \  }}||z  }t        |      D ]&  }|dddd|ddf   |dddd||z   ddf<   |dz  }( @ |S )zDConvert compact GroupPose keypoints to class-padded keypoint layout.rH   z8_format_keypoint_output received tensor with batch_size=z, num_queries=z but expected batch_size=zW. Shape mismatch silently bypassed in earlier versions; raise to surface upstream bugs.zE_format_keypoint_output received tensor with total_compact_keypoints=z) but schema expects either compact total=z or padded total=z for num_keypoints_per_class=.r   r   r   Nr!   )rm   rS   r;   r$   rn   r   rK   r   r,   r   r   r   r   rw   )r~   r   r   r   
batch_sizer[   total_compact_keypointskeypoint_dimmax_num_keypointstotal_padded_keypointstotal_actual_keypointspaddedcompact_idxr   keypoint_countclass_offsetkeypoint_idxs                    r5   _format_keypoint_outputzLWDETR._format_keypoint_outputZ  s    ++43O3O$$  "a'$$IZI`I`F
K!8,,,?S0SJ:,Vdepdq r++>*?~NbMc dhh    < <=!$T%A%A!BEV!V!$T%A%A!B"&<<$$"&<<WXoWp q;;Q:R S##9"::WX\XtXtWuuvx  "#))$++
 )243O3O)P 	!%I~$'88L %n 5 !?PQRTUWbdeQe?fq!\L8!;<q !	!
 rB   keypoint_predictionsc           	         | j                   sRt        j                  g |j                  dd | j                  j
                  |j                  |j                        S t        | j                         }t        | j                         } |d   j                  g |j                  dd || }|| j                  j                  |j                        z  }|j                  d      }| j                  j
                  }|dz
  }|j                  d   |k  r|j                  d   |k  rI| j                  s=t        j!                  d|j                  d   ||j                  d   |dz
         d	| _        t        j"                  | |j$                  g |j                  dd ||j                  d   z
   gd      }|S |j                  d   |kD  r	|d
d|f   }|S )zIAggregate keypoint class-logit contributions into detection-class logits.Nr   ).rI   rF   r   r!   zKeypoint class-logit boost has %d classes but detection head has %d foreground classes; zero-padding boost for classes %d..%d. Detection classes with no keypoint schema will receive zero boost. This warning is emitted once per model instance.T.)rS   r,   r   r$   r^   rM   r   r   rK   r   viewrX   r   r   ro   loggerwarningcatr:   )r~   r   num_keypoint_classesr   class_contribclass_boostdetection_num_classesforeground_num_classess           r5    _aggregate_keypoint_class_logitsz'LWDETR._aggregate_keypoint_class_logits  s   ++;;Q&,,Sb1Q43C3C3P3PQ*00+22   #4#?#?@ < <=9,V499 
!'',
.B
DU
 &(<(<(?(?@S@S(TT#''B'/ $ 0 0 = =!6!:R #88  $'==
 //NNd $))"-.#))"-.2 04D,)))K))q;+<+<Sb+AqCX[f[l[lmo[pCpq K  r"%:: &c+A,A+A&ABKrB   samplesc           	         t        |t        t        j                  f      rt	        |      }| j                  |      \  }}}g }g }|D ];  }|j                         \  }	}
|j                  |	       |j                  |
       |
;J  | j                  r-| j                  j                  }| j                  j                  }nF| j                  j                  d| j                   }| j                  j                  d| j                   }| j                  8| j                  r| j                  j                  n| j                  j                  }d}|-g }|D ]&  }|j                         \  }}|j                  |       ( | j!                  ||||||      }| j"                  r|\  }}}}}}}n|dd \  }}}}d}d}|b| j$                  ri| j'                  |      }|dddf   |dddf   z  |dddf   z   }|dddf   j)                         |dddf   z  }t        j*                  ||gd      }n"| j'                  |      |z   j-                         }| j/                  |      }d}| j"                  r| j0                  
|t3        d      | j1                  |      }|dddf   j5                  d	      } |dddf   j5                  d	      }!|dddf   | z  |!z   }"|dddf   }#t        j6                  |"|#gd      }$g }%t9        |$j:                  d
         D ]C  }&|$|&   }'|%j                  | j=                  |'|'j:                  d
   |'j:                  d                E t        j>                  |%d
      }|| jA                  |      z   }| j                  . |d
   jB                  ||jB                  j:                  d	d       }(|d   |d   d})| j                  (d   |)d<   ||d   |)d<   | jD                  r%| jG                  ||| j                  (nd|      |)d<   | jH                  ra| j                  r| jJ                  nd}*|jM                  |*d      }+g },t9        |*      D ]5  }- | j                   jN                  |-   |+|-         }.|,j                  |.       7 t        j6                  |,d      },d}/| j"                  rC|A| j=                  ||j:                  d
   |j:                  d         }/|,| jA                  |/      z   },| j                  4 |d
   jB                  |g|jB                  j:                  d	d d      d
   }0|(|,|d)d<   | j                  0|)d   d<   |/|/|)d   d<   |)S |,|d})| j                  0|)d<   |/|/|)d<   )S )a8  The forward expects a NestedTensor, which consists of:

           - samples.tensor: batched images, of shape [batch_size x 3 x H x W]
           - samples.mask: a binary mask of shape [batch_size x H x W], containing 1 on padded pixels

        It returns a dict with the following elements:
           - "pred_logits": the classification logits (including no-object) for all queries.
                            Shape= [batch_size x num_queries x num_classes]
           - "pred_boxes": The normalized boxes coordinates for all queries, represented as
                           (center_x, center_y, width, height). These values are normalized in [0, 1], relative to the
                           size of each individual image (disregarding possible padding). See PostProcess for
                           information on how to retrieve the unnormalized bounding box.
           - "aux_outputs": Optional, only returned when auxiliary losses are activated. It is a list of
                            dictionaries containing the two above keys for each decoder layer.
        Ncross_attn_srcsrH   .r   rF   r   Kuse_grouppose_keypoints=True requires keypoint_hs from transformer outputs.r   r   r!   )pred_logits
pred_boxes
pred_maskspred_keypointsaux_outputsTskip_blocksenc_outputs)(rL   r   r,   r   r   rg   	decomposeappendtrainingrb   r#   rc   r[   r`   sparse_forwardr   r\   rm   rl   r_   expconcatsigmoidr^   rp   rn   	unsqueezer   rw   r$   r   stackr   tensorsrh   _set_aux_lossru   ri   chunkr{   )1r~   r   targetsfeaturesposscross_attn_featuressrcsmasksfeatsrcr   refpoint_embed_weightquery_feat_weightseg_head_fwdr   feature	cross_srcr   transformer_outputshsref_unsigmoidhs_encref_enckeypoint_hsenc_kp_predictionsoutputs_coord_deltaoutputs_coord_cxcyoutputs_coord_whoutputs_coordoutputs_classoutputs_keypointsoutputs_keypoints_deltaref_whref_xykeypoints_xykeypoints_otheroutputs_keypoints_compactlayer_outputs_keypoints	layer_idxcompact_predsoutputs_masksoutri   hs_enc_listcls_encg_idxcls_enc_gidxkeypoints_enc	masks_encs1                                                    r5   r   zLWDETR.forward  sL     gell344W=G.2mmG.D+$+ 	$D(ICKKLL###		$ ==$($7$7$>$>! $ 6 6 %)$7$7$>$>?QAQAQ$R! $ 6 67I9I9I J!!-DHMM411@@W[WmWmWuWuL* O. 2&002	1&&y12 #..!+ / 
 ''UhRBvw=OQR1DRa1H.BvwK!%>  &*oob&9#%8bqb%AMRUWXWYRYDZ%Z]jknprqrprkr]s%s"#6sABw#?#C#C#EVY[\[]V]H^#^  %.@BR-SY[ \!%!4}!D M M O ,,R0M $++0C0C0O&$%rss*.*=*=k*J'&sABw/99"=&sBQBw/99"=6sBQBw?&H6Q"9#qr'"B,1II|_6U[],^)*,'!&'@'F'Fq'I!J I$=i$HM+2244))//2)//2 %*KK0GQ$O! -0U0UVg0h h%%1 ,Xa[-@-@"gooF[F[\^\_F` a"/"3=QSCTUC%%1$1"$5L! ,(9"(=$%}}%)%7%7!!%)%;%;%GMT%	&M" >>,0MMqJ ,,zq,9KGz* -Jt//CCEJ;W\K]^|,- iiQ/G M++0B0N $ < <&&,,Q/&,,Q/!
 "D$I$I-$XX%%1(QK'' OO))"#. $ 	 ~5<G%TM"))57@C&|4 ,;HC&'78 
 '.WE))5(1C% ,,9C()
rB   c                 l   | j                  |      \  }}}}| j                  j                  d | j                   }| j                  j                  d | j                   }| j                  |d ||||      }| j                  r|\  }	}
}}}}}n|d d \  }	}
}}d }d }d }d }|	| j                  ri| j                  |	      }|dd df   |
ddd f   z  |
dd df   z   }|ddd f   j                         |
ddd f   z  }t        j                  ||gd      }n"| j                  |	      |
z   j                         }| j                  |	      }| j                  r| j                  |t        d      | j                  |      }|
ddd f   j!                  d      }|
dd df   j!                  d      }|dd df   |z  |z   }|ddd f   }t        j"                  ||gd      }|j%                         d	k(  r|d   }| j'                  ||j(                  d
   |j(                  d         }|| j+                  |      z   }| j,                  | j-                  |d
   |	g|j(                  dd        d
   }n| j.                  sJ d        | j
                  j0                  d
   |      }|}| j                  rC|A| j'                  ||j(                  d
   |j(                  d         }|| j+                  |      z   }| j,                  )| j-                  |d
   |g|j(                  dd  d      d
   }||||fS ||||fS ||fS )Nr   rH   .r   rF   r   r   r      r   r!   z,if not using decoder, two_stage must be TrueTr   )rg   rb   r#   r[   rc   r\   rm   rl   r_   r   r,   r   r   r^   rp   rn   r   r   r;   r   r$   r   r`   ru   r{   )r~   r   r  r   r   r   r  r  r
  r  r  r  r  r  r  r   r  r  r  r  r  r  r  r  r  r  r  s                              r5   r   zLWDETR.forward_export[  s   )-w)?&a $ 3 3 : :;MT=M=M N OO223ET5E5EF"..!+ / 
 ''UhRBvw=OQR1DRa1H.BvwK!% >  &*oob&9#%8bqb%AMRUWXWYRYDZ%Z]jknprqrprkr]s%s"#6sABw#?#C#C#EVY[\[]V]H^#^  %.@BR-SY[ \!%!4}!D M M O ,,R0M++0C0C0O&$%rss*.*=*=k*J'&sABw/99"=&sBQBw/99"=6sBQBw?&H6Q"9#qr'"B$)II|_.MSU$V!$((*a/(9"(=%$($@$@%%++A.%++A.%!
 !.0U0UVg0h h%%1 $ 6 6G MM"#&! ! >>Q#QQ>CD,,@@CFKM#M++0B0N$($@$@&&,,Q/&,,Q/%!
 !.0U0UVg0h h%%1 $ 6 6G MM"#& $ !7 ! ! $ ->>( -1BBB -//rB   r  r  r   r  c                 D   ddg}|d d |d d g}|%|j                  d       |j                  |d d        |%|j                  d       |j                  |d d        t        | D 	cg c]"  }t        ||      D 	ci c]  \  }}	||	
 c}	}$ c}	}}S c c}	}w c c}	}}w )Nr   r   rF   r   r   )r   zip)
r~   r  r  r   r  namesvalueslayer_valuesr   values
             r5   r   zLWDETR._set_aux_loss  s     -$mCR&89$LL&MM-,-(LL)*MM+CR01^aci^jkklE<0HIuuIkkIks   )B=B
BBc                    | j                   d   j                  }t        |d      r|j                  S t        |d      r,t        |j                  d      r|j                  j                  S t        |d      rVt        |j                  d      r@t        |j                  j                  d      r |j                  j                  j
                  S y)a_  Resolve the list of transformer blocks/layers from backbone[0].encoder.

        Supports multiple backbone architectures:
        - encoder.blocks (standard ViT)
        - encoder.trunk.blocks (aimv2)
        - encoder.encoder.encoder.layer (HuggingFace DinoV2)

        Returns:
            List of transformer layers, or None if not found.
        r   blockstrunkencoderlayerN)rg   r3  r   r1  r2  r4  )r~   encs     r5   _get_backbone_encoder_layersz#LWDETR._get_backbone_encoder_layers  s     mmA&&3!::3 WSYY%A99###3	"ws{{I'F7SVS^S^SfSfhoKp;;&&,,,rB   drop_path_ratevit_encoder_num_layersc                 d   | j                         }|yt        |t        |            }t        j                  d||      D cg c]  }|j                          }}t        |      D ]C  }t        ||   d      st        ||   j                  d      s-||   ||   j                  _	        E yc c}w )a  Update drop_path rates for backbone encoder layers with linear schedule.

        Applies a linear schedule where the first layer has drop_path_rate=0 and the last layer has
        drop_path_rate=drop_path_rate. Intermediate layers are interpolated linearly.

        Args:
            drop_path_rate: Maximum drop path rate (applied to last layer).
            vit_encoder_num_layers: Number of encoder layers to update.
        Nr   	drop_path	drop_prob)
r6  minrK   r,   linspaceitemrw   r   r:  r;  )r~   r7  r8  rE   nxdp_ratesis           r5   update_drop_pathzLWDETR.update_drop_path  s     224>&F4&+nnQ&JKAFFHKKq 	<Avay+.76!9;N;NP[3\08q	##-	< Ls   B-c                     | j                   j                         D ]$  }t        |t        j                        s||_        & y N)r\   modulesrL   r   Dropoutp)r~   	drop_raterC   s      r5   update_dropoutzLWDETR.update_dropout  s5    &&..0 	%F&"**-$	%rB   )Fr!   FFFFNr!   )r   NrE  )$__name__
__module____qualname____doc__r   r%   rZ   r   staticmethodr,   r   rr   r   r   dictstrr   r   r   r   r   r   r   r   r   jitunusedr   r   r   rv   r6  floatrC  rJ  __classcell__)r   s   @r5   rR   rR   t   s3   J " %4801g "&cT!1g +.gRs t & 	S	 	ell 	 	 T#Y 5<<  (bT#Y b YS%,,EV@W Y\`ad\ehl\l Y YW$s)dBR WW[ W@E$1 <<1 !1 "	1
 
1f0U\\ 0V[VbVb 0dZ| ZxT0l YY 26l||l ||l ||d*	l
 !<<$.l l(hr}}.E (<u <c <VZ <&%rB   rR   c                 8   | j                   dz   }t        j                  | j                         t        d#i d| j                  d| j
                  d| j                  d| j                  d| j                  d| j                  d| j                  d	| j                  d
| j                  d| j                  d| j                  d| j                  d| j                  dt!        | d      r| j"                  n%t!        | d      r| j$                  | j$                  fndd| j&                  d| j(                  d| j*                  d| j,                  d| j.                  d u d| j0                  d| j2                  d| j4                  d| j6                  }| j8                  r|d   j                  d d fS | j:                  r|d d fS t=        | j                        | _        tA        |       }| jB                  r,tE        | j                  | jF                  | jH                        nd }tK        ||||| jL                  | jN                  | jP                  | jR                  | jT                  | jV                  tY        | dd      tY        | d g       tY        | d!d      "      }|S )$Nr!   r3  r8  pretrained_encoderwindow_block_indexesr:  out_channelsout_feature_indexesprojector_scaleuse_cls_tokenr   position_embeddingfreeze_encoder
layer_normtarget_shaper$   
resolution)  rb  rms_normbackbone_loraforce_no_pretraingradient_checkpointingload_dinov2_weights
patch_sizenum_windowspositional_encoding_sizedual_projectorr   )downsample_ratiorm   FrS   rT   )
r   r[   rh   ri   ru   rj   rl   rm   rS   rT    )-r   r,   r   r   r3  r8  rW  rX  r:  r   rZ  r[  r\  r]  r^  r_  r   r$   ra  rc  rd  re  rf  pretrain_weightsrh  ri  rj  rk  encoder_onlybackbone_onlyrK   num_feature_levelsr   r`   r   
dec_layersmask_downsample_ratiorR   r[   rh   ri   ru   rj   rl   rJ   )argsr   rg   r\   r`   models         r5   build_modelrv    s    ""Q&K	LL #::  22 "66	
 .. __ !44 ,, (( ??  22 ** ??  tW% JJ8?l8S4??DOO4Yc#& '( (()* 00+,  $::-. !11T9/0 ??12 $$34 "&!>!>56 **7H: {""D$..t##!$"6"67D#D)K !! 	OOOO!77	
   $$??..!66&& '.G O '.G L)07Y[\)]E  LrB   c                    t        j                  | j                        }t        |       }| j                  | j                  d}| j
                  |d<   | j                  r| j                  |d<   | j                  |d<   t        | dd      }|r@t        | dd      |d	<   t        | d
d      |d<   t        | dd      |d<   t        | dd      |d<   | j                  ri }t        | j                  dz
        D ];  }|j                  |j                         D ci c]  \  }}|d| z   | c}}       = | j                  r6|j                  |j                         D ci c]  \  }}|dz   | c}}       |j                  |       g d}	| j                  r|	j!                  d       |r|	j!                  d       t        | dd      }
| j                  rlt#        | j$                  dz   ||| j&                  |	| j(                  |
| j*                  | j,                  | j.                  | j0                  t        | dg             }n`t#        | j$                  dz   ||| j&                  |	| j(                  |
| j*                  | j,                  | j.                  t        | dg             }|j3                  |       t5        | j6                  t        | dg       t        | dd            }||fS c c}}w c c}}w )N)loss_ce	loss_bbox	loss_giouloss_mask_celoss_mask_dicerm   Fkeypoint_l1_loss_coefg        loss_keypoints_l1keypoint_findable_loss_coefloss_keypoints_findablekeypoint_visible_loss_coefloss_keypoints_visiblekeypoint_nll_loss_coefloss_keypoints_nllr!   r   _enc)labelsboxescardinalityr  	keypointssum_group_lossesrS   )matcherweight_dictfocal_alphalossesri   r  use_varifocal_lossuse_position_supervised_lossia_bce_lossmask_point_sample_ratiorS   )
r  r  r  r  ri   r  r  r  r  rS   postprocess_trace_alphag?)
num_selectrS   trace_alpha)r,   r   r   cls_loss_coefbbox_loss_coefgiou_loss_coefr`   mask_ce_loss_coefmask_dice_loss_coefrJ   rh   rw   rr  updateitemsru   r   r   r   r  ri   r  r  r  r  r   r   r  )rt  r   r  r  has_keypointsaux_weight_dictrB  kvr  r  	criterionpostprocesss                r5   "build_criterion_and_postprocessorsr  A  s   \\$++&FD!G"00t?R?RSK#22K&*&<&<N#(,(@(@$%D";UCM+249PRU+V'(18?\^a1b-.07>Z\_0`,-,3D:RTW,X()}}t*+ 	UA""{?P?P?R#Stq!A!A3KN#ST	U>>""k>O>O>Q#RdaAJM#RS?+/Fgk"t%7? q #((-#66)-)J)J(($($@$@$+D2KR$P
	 !q #((-#66)-)J)J(($+D2KR$P
	 LL?? '.G LD";SA	K k!!c $T#Rs   KK%
model_configr   train_configr   defaultsc                     ddl m} |j                  s|j                  rt	        d      |ddlm}  |dd      } || ||      }t        |      S )uQ  Build an LWDETR model directly from a ModelConfig.

    A config-native alternative to ``build_model(build_namespace(mc, tc))``. Constructs the namespace internally from
    ``model_config``, an optional ``train_config``, and ``defaults``, then delegates to :func:`build_model`.

    Note:
        The internal ``SimpleNamespace`` bridge is transitional — it will be eliminated once all builder functions
        accept config objects directly. Callers should not rely on the namespace shape or pass it externally.

    Args:
        model_config: Architecture configuration.
        train_config: Training hyperparameter configuration. If ``None``,
            a minimal dummy ``TrainConfig(dataset_dir=".", output_dir=".")`` is constructed, matching the previous
            default behavior.
        defaults: Hardcoded architectural constants. Defaults to ``MODEL_DEFAULTS``.

    Returns:
        Fully initialised LWDETR model.

    Raises:
        ValueError: If ``defaults`` request ``encoder_only`` or ``backbone_only``,
            which would make the return type differ from ``LWDETR``.
    r   _namespace_from_configsz`build_model_from_config() requires defaults.encoder_only=False and defaults.backbone_only=False.)r   r   )dataset_dir
output_dir)rfdetr._namespacer  ro  rp  rn   rfdetr.configr   rv  )r  r  r  r  r   nss         r5   build_model_from_configr    sW    8 : 6 6n
 	
 -"ssC	 |X	FBr?rB   c                 8    ddl m}  || ||      }t        |      S )a  Build criterion and postprocessor directly from config objects.

    A config-native alternative to ``build_criterion_and_postprocessors(build_namespace(mc, tc))``.

    Args:
        model_config: Architecture configuration.
        train_config: Training hyperparameter configuration.
        defaults: Hardcoded architectural constants. Defaults to ``MODEL_DEFAULTS``.

    Returns:
        A 2-tuple of ``(SetCriterion, PostProcess)``.
    r   r  )r  r  r  )r  r  r  r  r  s        r5   build_criterion_from_configr    s     " :	 |X	FB-b11rB   )rt  r   )9rN  rx   r&   typingr   r   r   r,   r   r  r   r   rfdetr.models._defaultsr	   r
   rfdetr.models._typesr   rfdetr.models.backboner   rfdetr.models.criterionr   r   r   r   r   r   r   r    rfdetr.models.heads.segmentationr   rfdetr.models.matcherr   rfdetr.models.mathr   rfdetr.models.postprocessr   rfdetr.models.transformerr   rfdetr.utilities.loggerr   rfdetr.utilities.tensorsr   r   r   r*   r%   r6   r<   rA   ModulerP   rR   rv  r  r  tupler  rm  rB   r5   <module>r     sG  $ +   4 4  6 A , 1	 	 	 > / " 1 7 . Q	299 3 299 <Qbll Qc Qbll Q* *t * |	%RYY |	%~K\D"R -1,))=)) ) 	)^ -222 2 <$%	2rB   