Ë
    þÍ:jÊ=  ã                   ó  — d dl mZ d dlZd dlmZmZ d dlmZ d dlm	Z
mZ ddlmZ ddlmZ dd	lmZ  G d
„ dej&                  «      Zdededededededefd„Zdee   dee   deeef   fd„Z G d„ dej                  j&                  «      Zy)é    )ÚOptionalN)ÚnnÚTensor)Ú
functional)ÚboxesÚConv2dNormActivationé   )Ú_utils)ÚAnchorGenerator)Ú	ImageListc                   ól   ‡ — e Zd ZdZdZddededdfˆ fd„Zˆ fd„Zd	ee	   de
ee	   ee	   f   fd
„Zˆ xZS )ÚRPNHeada  
    Adds a simple RPN Head with classification and regression heads

    Args:
        in_channels (int): number of channels of the input feature
        num_anchors (int): number of anchors to be predicted
        conv_depth (int, optional): number of convolutions
    é   Úin_channelsÚnum_anchorsÚreturnNc           	      óz  •— t         ‰| �  «        g }t        |«      D ]   }|j                  t	        ||dd ¬«      «       Œ" t        j                  |Ž | _        t        j                  ||dd¬«      | _	        t        j                  ||dz  dd¬«      | _
        | j                  «       D ]“  }t        |t
        j                  «      sŒt        j
                  j                  j                  |j                   d¬«       |j"                  €Œ`t        j
                  j                  j%                  |j"                  d«       Œ• y )	Né   )Úkernel_sizeÚ
norm_layerr	   )r   Ústrideé   g{®Gáz„?)Ústdr   )ÚsuperÚ__init__ÚrangeÚappendr   r   Ú
SequentialÚconvÚConv2dÚ
cls_logitsÚ	bbox_predÚmodulesÚ
isinstanceÚtorchÚinitÚnormal_ÚweightÚbiasÚ	constant_)Úselfr   r   Ú
conv_depthÚconvsÚ_ÚlayerÚ	__class__s          €úu/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchvision/models/detection/rpn.pyr   zRPNHead.__init__   sí   ø€ Ü‰ÑÔØˆÜ�zÓ"ò 	iˆAØ�L‰LÔ-¨k¸;ÐTUÐbfÔgÕhð	iä—M‘M 5Ð)ˆŒ	ÜŸ)™) K°È!ÐTUÔVˆŒÜŸ™ ;°¸a±ÈQÐWXÔYˆŒà—\‘\“^ò 	;ˆEÜ˜%¤§¡Õ+Ü—‘—‘×%Ñ% e§l¡l¸Ð%Ô=Ø—:‘:Ñ)Ü—H‘H—M‘M×+Ñ+¨E¯J©J¸Õ:ñ		;ó    c           	      ó¾   •— |j                  dd «      }|�|dk  r.dD ])  }	|› d|	› �}
|› d|	› �}|
|v sŒ|j                  |
«      ||<   Œ+ t        ‰| �  |||||||«       y )NÚversionr   )r(   r)   zconv.z	conv.0.0.)ÚgetÚpopr   Ú_load_from_state_dict)r+   Ú
state_dictÚprefixÚlocal_metadataÚstrictÚmissing_keysÚunexpected_keysÚ
error_msgsr4   ÚtypeÚold_keyÚnew_keyr0   s               €r1   r7   zRPNHead._load_from_state_dict*   s’   ø€ ð !×$Ñ$ Y°Ó5ˆàˆ?˜g¨škØ*ò B�Ø#˜H E¨$¨Ð0�Ø#˜H I¨d¨VÐ4�Ø˜jÒ(Ø*4¯.©.¸Ó*A�J˜wÒ'ð	Bô 	‰Ñ%ØØØØØØØõ	
r2   Úxc                 óÂ   — g }g }|D ]S  }| j                  |«      }|j                  | j                  |«      «       |j                  | j                  |«      «       ŒU ||fS ©N)r   r   r!   r"   )r+   rB   ÚlogitsÚbbox_regÚfeatureÚts         r1   ÚforwardzRPNHead.forwardG   s`   € ØˆØˆØò 	/ˆGØ—	‘	˜'Ó"ˆAØ�M‰M˜$Ÿ/™/¨!Ó,Ô-Ø�O‰O˜DŸN™N¨1Ó-Õ.ð	/ð �xÐÐr2   )r	   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ú_versionÚintr   r7   Úlistr   ÚtuplerI   Ú__classcell__©r0   s   @r1   r   r      sW   ø„ ñð €Hñ; Cð ;°cð ;ÈDõ ;ô
ð: ˜˜f™ð  ¨%°°V±¸dÀ6¹lÐ0JÑ*K÷  r2   r   r/   ÚNÚAÚCÚHÚWr   c                 ó€   — | j                  |d|||«      } | j                  ddddd«      } | j                  |d|«      } | S )Néÿÿÿÿr   r   r   r	   r   )ÚviewÚpermuteÚreshape)r/   rT   rU   rV   rW   rX   s         r1   Úpermute_and_flattenr^   Q   sE   € Ø�J‰J�q˜"˜a  AÓ&€EØ�M‰M˜!˜Q  1 aÓ(€EØ�M‰M˜!˜R Ó#€EØ€Lr2   Úbox_clsÚbox_regressionc           	      ó®  — g }g }t        | |«      D ]q  \  }}|j                  \  }}}}	|j                  d   }
|
dz  }||z  }t        ||||||	«      }|j                  |«       t        |||d||	«      }|j                  |«       Œs t	        j
                  |d¬«      j                  dd«      } t	        j
                  |d¬«      j                  dd«      }| |fS )Nr	   r   ©Údimr   éþÿÿÿrZ   )ÚzipÚshaper^   r   r%   ÚcatÚflattenr]   )r_   r`   Úbox_cls_flattenedÚbox_regression_flattenedÚbox_cls_per_levelÚbox_regression_per_levelrT   ÚAxCrW   rX   ÚAx4rU   rV   s                r1   Úconcat_box_prediction_layersro   X   sü   € ØÐØ!Ðô
 8;¸7ÀNÓ7Sò 	BÑ3ÐÐ3Ø(×.Ñ.‰ˆˆ3��1Ø&×,Ñ,¨QÑ/ˆØ�1‰HˆØ�1‰HˆÜ/Ð0AÀ1ÀaÈÈAÈqÓQÐØ× Ñ Ð!2Ô3ä#6Ð7OÐQRÐTUÐWXÐZ[Ð]^Ó#_Ð Ø ×'Ñ'Ð(@ÕAð	Bô �i‰iÐ)¨qÔ1×9Ñ9¸!¸RÓ@€GÜ—Y‘YÐ7¸QÔ?×GÑGÈÈAÓN€NØ�NÐ"Ð"r2   c                   ó  ‡ — e Zd ZdZej
                  ej                  ej                  dœZ	 d"de	de
j                  dedededed	eeef   d
eeef   dededdfˆ fd„Zdefd„Zdefd„Zdee   deeeef      deee   ee   f   fd„Zdedee   defd„Zdededeeeef      dee   deee   ee   f   f
d„Zdededee   dee   deeef   f
d„Z	 d#ded eeef   deeeeef         deee   eeef   f   fd!„Zˆ xZS )$ÚRegionProposalNetworkaÏ  
    Implements Region Proposal Network (RPN).

    Args:
        anchor_generator (AnchorGenerator): module that generates the anchors for a set of feature
            maps.
        head (nn.Module): module that computes the objectness and regression deltas
        fg_iou_thresh (float): minimum IoU between the anchor and the GT box so that they can be
            considered as positive during training of the RPN.
        bg_iou_thresh (float): maximum IoU between the anchor and the GT box so that they can be
            considered as negative during training of the RPN.
        batch_size_per_image (int): number of anchors that are sampled during training of the RPN
            for computing the loss
        positive_fraction (float): proportion of positive anchors in a mini-batch during training
            of the RPN
        pre_nms_top_n (Dict[str, int]): number of proposals to keep before applying NMS. It should
            contain two fields: training and testing, to allow for different values depending
            on training or evaluation
        post_nms_top_n (Dict[str, int]): number of proposals to keep after applying NMS. It should
            contain two fields: training and testing, to allow for different values depending
            on training or evaluation
        nms_thresh (float): NMS threshold used for postprocessing the RPN proposals
        score_thresh (float): only return proposals with an objectness score greater than score_thresh

    )Ú	box_coderÚproposal_matcherÚfg_bg_samplerÚanchor_generatorÚheadÚfg_iou_threshÚbg_iou_threshÚbatch_size_per_imageÚpositive_fractionÚpre_nms_top_nÚpost_nms_top_nÚ
nms_threshÚscore_threshr   Nc                 óT  •— t         ‰| �  «        || _        || _        t	        j
                  d¬«      | _        t        j                  | _	        t	        j                  ||d¬«      | _        t	        j                  ||«      | _        || _        || _        |	| _        |
| _        d| _        y )N)ç      ð?r€   r€   r€   )ÚweightsT)Úallow_low_quality_matchesgü©ñÒMbP?)r   r   ru   rv   Ú	det_utilsÚBoxCoderrr   Úbox_opsÚbox_iouÚbox_similarityÚMatcherrs   ÚBalancedPositiveNegativeSamplerrt   Ú_pre_nms_top_nÚ_post_nms_top_nr}   r~   Úmin_size)r+   ru   rv   rw   rx   ry   rz   r{   r|   r}   r~   r0   s              €r1   r   zRegionProposalNetwork.__init__’   s›   ø€ ô 	‰ÑÔØ 0ˆÔØˆŒ	Ü"×+Ñ+Ð4HÔIˆŒô &Ÿo™oˆÔä )× 1Ñ 1ØØØ&*ô!
ˆÔô '×FÑFÐG[Ð]nÓoˆÔà+ˆÔØ-ˆÔØ$ˆŒØ(ˆÔØˆ�r2   c                 óV   — | j                   r| j                  d   S | j                  d   S ©NÚtrainingÚtesting)r�   rŠ   ©r+   s    r1   r{   z#RegionProposalNetwork.pre_nms_top_n·   s+   € Ø�=Š=Ø×&Ñ& zÑ2Ð2Ø×"Ñ" 9Ñ-Ð-r2   c                 óV   — | j                   r| j                  d   S | j                  d   S rŽ   )r�   r‹   r‘   s    r1   r|   z$RegionProposalNetwork.post_nms_top_n¼   s+   € Ø�=Š=Ø×'Ñ'¨
Ñ3Ð3Ø×#Ñ# IÑ.Ð.r2   ÚanchorsÚtargetsc                 óÆ  — g }g }t        ||«      D �]J  \  }}|d   }|j                  «       dk(  rq|j                  }t        j                  |j
                  t        j                  |¬«      }	t        j                  |j
                  d   ft        j                  |¬«      }
n™| j                  ||«      }| j                  |«      }||j                  d¬«         }	|dk\  }
|
j                  t        j                  ¬«      }
|| j                  j                  k(  }d|
|<   || j                  j                  k(  }d|
|<   |j                  |
«       |j                  |	«       �ŒM ||fS )Nr   r   ©ÚdtypeÚdevice)Úmin)r—   ç        g      ð¿)re   Únumelr˜   r%   Úzerosrf   Úfloat32r‡   rs   ÚclampÚtoÚBELOW_LOW_THRESHOLDÚBETWEEN_THRESHOLDSr   )r+   r“   r”   ÚlabelsÚmatched_gt_boxesÚanchors_per_imageÚtargets_per_imageÚgt_boxesr˜   Úmatched_gt_boxes_per_imageÚlabels_per_imageÚmatch_quality_matrixÚmatched_idxsÚ
bg_indicesÚinds_to_discards                  r1   Úassign_targets_to_anchorsz/RegionProposalNetwork.assign_targets_to_anchorsÁ   sq  € ð ˆØÐÜ47¸ÀÓ4Ió 	@Ñ0ÐÐ0Ø(¨Ñ1ˆHà�~‰~Ó 1Ò$à*×1Ñ1�Ü-2¯[©[Ð9J×9PÑ9PÔX]×XeÑXeÐntÔ-uÐ*Ü#(§;¡;Ð0A×0GÑ0GÈÑ0JÐ/LÔTY×TaÑTaÐjpÔ#qÑ à'+×':Ñ':¸8ÐEVÓ'WÐ$Ø#×4Ñ4Ð5IÓJ�ð
 .6°l×6HÑ6HÈQÐ6HÓ6OÑ-PÐ*à#/°1Ñ#4Ð Ø#3×#6Ñ#6¼U¿]¹]Ð#6Ó#KÐ ð *¨T×-BÑ-B×-VÑ-VÑV�
Ø/2Ð  Ñ,ð #/°$×2GÑ2G×2ZÑ2ZÑ"Z�Ø48Ð  Ñ1à�M‰MÐ*Ô+Ø×#Ñ#Ð$>Ö?ð;	@ð< Ð'Ð'Ð'r2   Ú
objectnessÚnum_anchors_per_levelc                 ó,  — g }d}|j                  |d«      D ]e  }|j                  d   }t        j                  || j	                  «       d«      }|j                  |d¬«      \  }}	|j                  |	|z   «       ||z  }Œg t        j                  |d¬«      S )Nr   r	   rb   )	Úsplitrf   rƒ   Ú	_topk_minr{   Útopkr   r%   rg   )
r+   r®   r¯   ÚrÚoffsetÚobr   r{   r.   Ú	top_n_idxs
             r1   Ú_get_top_n_idxz$RegionProposalNetwork._get_top_n_idxç   s—   € ØˆØˆØ×"Ñ"Ð#8¸!Ó<ò 	"ˆBØŸ(™( 1™+ˆKÜ%×/Ñ/°°D×4FÑ4FÓ4HÈ!ÓLˆMØŸ7™7 =°a˜7Ó8‰LˆAˆyØ�H‰H�Y Ñ'Ô(Ø�kÑ!‰Fð	"ô �y‰y˜ Ô"Ð"r2   Ú	proposalsÚimage_shapesc           
      ó(  — |j                   d   }|j                  }|j                  «       }|j                  |d«      }t	        |«      D ��cg c]-  \  }}t        j                  |f|t
        j                  |¬«      ‘Œ/ }	}}t        j                  |	d«      }	|	j                  dd«      j                  |«      }	| j                  ||«      }
t        j                  ||¬«      }|d d …d f   }|||
f   }|	||
f   }	|||
f   }t        j                  |«      }g }g }t        |||	|«      D ]á  \  }}}}t        j                  ||«      }t        j                   || j"                  «      }||   ||   ||   }}}t        j$                  || j&                  k\  «      d   }||   ||   ||   }}}t        j(                  |||| j*                  «      }|d | j-                  «        }||   ||   }}|j/                  |«       |j/                  |«       Œã ||fS c c}}w )Nr   rZ   r–   r	   )r˜   )rf   r˜   Údetachr]   Ú	enumerater%   ÚfullÚint64rg   Ú	expand_asr¸   ÚarangeÚsigmoidre   r…   Úclip_boxes_to_imageÚremove_small_boxesrŒ   Úwherer~   Úbatched_nmsr}   r|   r   )r+   r¹   r®   rº   r¯   Ú
num_imagesr˜   ÚidxÚnÚlevelsr·   Úimage_rangeÚ	batch_idxÚobjectness_probÚfinal_boxesÚfinal_scoresr   ÚscoresÚlvlÚ	img_shapeÚkeeps                        r1   Úfilter_proposalsz&RegionProposalNetwork.filter_proposalsò   s3  € ð —_‘_ QÑ'ˆ
Ø×!Ñ!ˆà×&Ñ&Ó(ˆ
Ø×'Ñ'¨
°BÓ7ˆ
ô S\Ð\qÓRr÷
ÙHNÈÈQŒE�J‰J˜�t˜S¬¯©¸FÖCð
ˆñ 
ô —‘˜6 1Ó%ˆØ—‘  2Ó&×0Ñ0°Ó<ˆð ×'Ñ'¨
Ð4IÓJˆ	ä—l‘l :°fÔ=ˆØ¢ 4 Ñ(ˆ	à 	¨9Ð 4Ñ5ˆ
Ø˜	 9Ð,Ñ-ˆØ˜i¨Ð2Ñ3ˆ	äŸ-™-¨
Ó3ˆàˆØˆÜ-0°¸OÈVÐUaÓ-bò 	(Ñ)ˆE�6˜3 	Ü×/Ñ/°°yÓAˆEô ×-Ñ-¨e°T·]±]ÓCˆDØ!& t¡¨f°T©l¸CÀ¹I˜3�6ˆEô —;‘;˜v¨×):Ñ):Ñ:Ó;¸AÑ>ˆDØ!& t¡¨f°T©l¸CÀ¹I˜3�6ˆEô ×&Ñ& u¨f°c¸4¿?¹?ÓKˆDð Ð/˜$×-Ñ-Ó/Ð0ˆDØ! $™K¨°©�6ˆEà×Ñ˜uÔ%Ø×Ñ Õ'ð)	(ð* ˜LÐ(Ð(ùóS
s   Á2HÚpred_bbox_deltasr¢   Úregression_targetsc                 ó,  — | j                  |«      \  }}t        j                  t        j                  |d¬«      «      d   }t        j                  t        j                  |d¬«      «      d   }t        j                  ||gd¬«      }|j	                  «       }t        j                  |d¬«      }t        j                  |d¬«      }t        j                  ||   ||   dd¬«      |j                  «       z  }t        j                  ||   ||   «      }	|	|fS )a  
        Args:
            objectness (Tensor)
            pred_bbox_deltas (Tensor)
            labels (List[Tensor])
            regression_targets (List[Tensor])

        Returns:
            objectness_loss (Tensor)
            box_loss (Tensor)
        r   rb   gÇqÇq¼?Úsum)ÚbetaÚ	reduction)	rt   r%   rÅ   rg   rh   ÚFÚsmooth_l1_lossr›   Ú binary_cross_entropy_with_logits)
r+   r®   rÕ   r¢   rÖ   Úsampled_pos_indsÚsampled_neg_indsÚsampled_indsÚbox_lossÚobjectness_losss
             r1   Úcompute_lossz"RegionProposalNetwork.compute_loss+  s  € ð .2×-?Ñ-?ÀÓ-GÑ*ÐÐ*Ü Ÿ;™;¤u§y¡yÐ1AÀqÔ'IÓJÈ1ÑMÐÜ Ÿ;™;¤u§y¡yÐ1AÀqÔ'IÓJÈ1ÑMÐä—y‘yÐ"2Ð4DÐ!EÈ1ÔMˆà×'Ñ'Ó)ˆ
ä—‘˜6 qÔ)ˆÜ"ŸY™YÐ'9¸qÔAÐä×#Ñ#ØÐ-Ñ.ØÐ/Ñ0ØØô	
ð
 ×ÑÓ!ñ#ˆô ×<Ñ<¸ZÈÑ=UÐW]Ð^jÑWkÓlˆà Ð(Ð(r2   ÚimagesÚfeaturesc                 óÐ  — t        |j                  «       «      }| j                  |«      \  }}| j                  ||«      }t	        |«      }|D �cg c]  }|d   j
                  ‘Œ }	}|	D �
cg c]  }
|
d   |
d   z  |
d   z  ‘Œ }}
t        ||«      \  }}| j                  j                  |j                  «       |«      }|j                  |dd«      }| j                  |||j                  |«      \  }}i }| j                  rZ|€t        d«      ‚| j                  ||«      \  }}| j                  j!                  ||«      }| j#                  ||||«      \  }}||dœ}||fS c c}w c c}
w )a=  
        Args:
            images (ImageList): images for which we want to compute the predictions
            features (Dict[str, Tensor]): features computed from the images that are
                used for computing the predictions. Each tensor in the list
                correspond to different feature levels
            targets (List[Dict[str, Tensor]]): ground-truth boxes present in the image (optional).
                If provided, each element in the dict should contain a field `boxes`,
                with the locations of the ground-truth boxes.

        Returns:
            boxes (List[Tensor]): the predicted boxes from the RPN, one Tensor per
                image.
            losses (Dict[str, Tensor]): the losses for the model during training. During
                testing, it is an empty dict.
        r   r	   r   rZ   r   ztargets should not be None)Úloss_objectnessÚloss_rpn_box_reg)rP   Úvaluesrv   ru   Úlenrf   ro   rr   Údecoder¼   r[   rÔ   Úimage_sizesr�   Ú
ValueErrorr­   Úencoderã   )r+   rä   rå   r”   r®   rÕ   r“   rÇ   ÚoÚ#num_anchors_per_level_shape_tensorsÚsr¯   r¹   r   rÐ   Úlossesr¢   r£   rÖ   rç   rè   s                        r1   rI   zRegionProposalNetwork.forwardP  sƒ  € ô. ˜Ÿ™Ó)Ó*ˆØ'+§y¡y°Ó':Ñ$ˆ
Ð$Ø×'Ñ'¨°Ó9ˆä˜“\ˆ
ØCMÖ.N¸a¨q°©t¯z«zÐ.NÐ+Ð.NØ=`Ö a¸  1¡¨¨!©¡¨q°©tÓ!3Ð aÐÐ aÜ'CÀJÐP`Ó'aÑ$ˆ
Ð$ð —N‘N×)Ñ)Ð*:×*AÑ*AÓ*CÀWÓMˆ	Ø—N‘N :¨r°1Ó5ˆ	Ø×-Ñ-¨i¸ÀV×EWÑEWÐYnÓo‰ˆˆvàˆØ�=Š=ØˆÜ Ð!=Ó>Ð>Ø'+×'EÑ'EÀgÈwÓ'WÑ$ˆFÐ$Ø!%§¡×!6Ñ!6Ð7GÈÓ!QÐØ04×0AÑ0AØÐ,¨fÐ6Hó1Ñ-ˆOÐ-ð $3Ø$4ñˆFð �fˆ}Ðùò/ /OùÚ as   ÁEÁ+E#)rš   rD   )rJ   rK   rL   rM   rƒ   r„   rˆ   r‰   Ú__annotations__r   r   ÚModuleÚfloatrO   ÚdictÚstrr   r{   r|   rP   r   rQ   r­   r¸   rÔ   rã   r   r   rI   rR   rS   s   @r1   rq   rq   q   s-  ø„ ñð6 ×'Ñ'Ø%×-Ñ-Ø"×BÑBñ€Oð& "ñ#à)ð#ð �i‰ið#ð
 ð#ð ð#ð "ð#ð !ð#ð ˜C ˜H‘~ð#ð ˜S #˜X™ð#ð ð#ð ð#ð 
õ#ðJ.˜só .ð
/ ó /ð
$(Ø˜F‘|ð$(Ø.2°4¸¸V¸Ñ3DÑ.Eð$(à	ˆt�F‰|˜T &™\Ð)Ñ	*ó$(ðL	#¨ð 	#ÈÈSÉ	ð 	#ÐV\ó 	#ð7)àð7)ð ð7)ð ˜5  c ™?Ñ+ð	7)ð
  $ C™yð7)ð 
ˆt�F‰|˜T &™\Ð)Ñ	*ó7)ðr#)Ø ð#)Ø4:ð#)ØDHÈÁLð#)ØfjÐkqÑfrð#)à	ˆv�vˆ~Ñ	ó#)ðR 6:ñ	3àð3ð �s˜F�{Ñ#ð3ð ˜$˜t C¨ KÑ0Ñ1Ñ2ð	3ð
 
ˆt�F‰|˜T # v +Ñ.Ð.Ñ	/÷3r2   rq   )Útypingr   r%   r   r   Útorch.nnr   rÛ   Útorchvision.opsr   r…   r   Ú r
   rƒ   Úanchor_utilsr   Ú
image_listr   rô   r   rO   r^   rP   rQ   ro   rq   © r2   r1   ú<module>rÿ      s¬   ðÝ ã ß Ý $ß Bå !õ *Ý !ô? ˆb�i‰iô ? ðD˜vð ¨#ð °#ð ¸#ð À#ð È#ð ÐRXó ð#¨$¨v©,ð #ÈÈVÉð #ÐY^Ð_eÐgmÐ_mÑYnó #ô2R˜EŸH™HŸO™Oõ Rr2   