§
    ‚ŠtjQ’ ã                   ó(  — d dl Z d dlmZm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  e¦   «         rd dl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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* ddl+m,Z,m-Z-m.Z. ddl/m0Z0 ddl1m2Z2 ddl3m4Z4 ddl5m6Z6m7Z7m8Z8m9Z9m:Z:m;Z;m<Z<  e*j=        e>¦  «        Z?e(e G d„ de¦  «        ¦   «         ¦   «         Z@e(e G d„ de ¦  «        ¦   «         ¦   «         ZAe(e G d„ de ¦  «        ¦   «         ¦   «         ZBe(e G d„ de ¦  «        ¦   «         ¦   «         ZCe(e G d „ d!e ¦  «        ¦   «         ¦   «         ZDe(e G d"„ d#e ¦  «        ¦   «         ¦   «         ZEdxd%ej        d&eFd'ej        fd(„ZGdyd*eHfd+„ZId,„ ZJ G d-„ d.e
jK        ¦  «        ZL	 	 dzd0e
jK        d1ej        d2ej        d3ej        d4ej        dz  d5eFdz  d6eFd7e%e,         fd8„ZM G d9„ d:e
jK        ¦  «        ZN G d;„ d<e
jK        ¦  «        ZOd=„ ZPd>ej        d?ej        d@ej        dAej        d'eQej        ej        f         f
dB„ZR G dC„ dDe
jK        ¦  «        ZS G dE„ dFe
jK        ¦  «        ZT G dG„ dHe
jK        ¦  «        ZUdI„ ZVdJ„ ZW G dK„ dLe
jK        ¦  «        ZX G dM„ dNe¦  «        ZYe( e0dO¬P¦  «         G dQ„ dRe#¦  «        ¦   «         ¦   «         ZZe( G dS„ dTeZ¦  «        ¦   «         Z[ G dU„ dVe
jK        ¦  «        Z\ G dW„ dXe
jK        ¦  «        Z] G dY„ dZe
jK        ¦  «        Z^ e(d[¬\¦  «         G d]„ d^eZ¦  «        ¦   «         Z_ G d_„ d`e
jK        ¦  «        Z` G da„ dbe
jK        ¦  «        Za G dc„ dde
jK        ¦  «        Zb G de„ dfeZ¦  «        Zc G dg„ dhe
jK        ¦  «        Zd G di„ dje
jK        ¦  «        Ze G dk„ dleZ¦  «        Zf G dm„ dne
jK        ¦  «        Zg G do„ dpe
jK        ¦  «        Zh G dq„ dre
jK        ¦  «        Zi G ds„ dteZ¦  «        Zj G du„ dveZ¦  «        Zkg dw¢ZldS ){é    N)ÚCallableÚIterable)Ú	dataclass)ÚTensoré   )Úis_torchvision_available)ÚCLIPTextModelWithProjection)Úinitialization)ÚACT2FN)Úcreate_bidirectional_mask)ÚGradientCheckpointingLayer)ÚBaseModelOutputÚBaseModelOutputWithPoolingÚModelOutput)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)Ú#compile_compatible_method_lru_cache)Úauto_docstringÚcan_return_tupleÚlogging)ÚTransformersKwargsÚis_flash_attention_requestedÚmerge_with_config_defaults)Úrequires)Úcapture_outputsé   )Ú	AutoModelé   )Ú
Sam3ConfigÚSam3DETRDecoderConfigÚSam3DETREncoderConfigÚSam3GeometryEncoderConfigÚSam3MaskDecoderConfigÚSam3VisionConfigÚSam3ViTConfigc                   ód   — e Zd ZU dZdZeej        df         ed<   dZ	eej        df         ed<   dS )ÚSam3VisionEncoderOutputzØ
    fpn_hidden_states (`tuple[torch.FloatTensor]`):
        Tuple of multi-level FPN feature maps.
    fpn_position_encoding (`tuple[torch.FloatTensor]`):
        Tuple of position encodings for each FPN level.
    N.Úfpn_hidden_statesÚfpn_position_encoding)
Ú__name__Ú
__module__Ú__qualname__Ú__doc__r)   ÚtupleÚtorchÚFloatTensorÚ__annotations__r*   © ó    úd/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/sam3/modeling_sam3.pyr(   r(   D   sZ   € € € € € € ðð ð 8<Ð�u˜UÔ.°Ð3Ô4Ð;Ð;Ñ;Ø;?Ð˜5 Ô!2°CÐ!7Ô8Ð?Ð?Ñ?Ð?Ð?r4   r(   c                   óJ   — e Zd ZU dZdZej        ed<   dZej	        dz  ed<   dS )ÚSam3GeometryEncoderOutputa^  
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_prompts, hidden_size)`):
        Encoded geometry prompt features (boxes).
    attention_mask (`torch.BoolTensor` of shape `(batch_size, num_prompts)`, *optional*):
        Attention mask for geometry prompts where True indicates valid positions and False indicates padding.
    NÚlast_hidden_stateÚattention_mask)
r+   r,   r-   r.   r8   r0   r1   r2   r9   Ú
BoolTensorr3   r4   r5   r7   r7   R   sJ   € € € € € € ðð ð ,0Ð�uÔ(Ð/Ð/Ñ/Ø.2€N�EÔ$ tÑ+Ð2Ð2Ñ2Ð2Ð2r4   r7   c                   óÚ   — e Zd ZU dZdZej        ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
ej        dz  ed<   dZeej                 dz  ed<   dZeej                 dz  ed<   dS )	ÚSam3DETREncoderOutputaˆ  
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
        Encoded vision features (flattened from multi-level features).
    pos_embeds_flattened (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
        Flattened position embeddings for the vision features.
    text_features (`torch.FloatTensor` of shape `(batch_size, text_seq_len, hidden_size)`, *optional*):
        Text features (may be pooled after encoder processing).
    spatial_shapes (`torch.LongTensor` of shape `(num_levels, 2)`, *optional*):
        Spatial shapes (height, width) for each feature pyramid level.
    hidden_states (`tuple[torch.FloatTensor]`, *optional*):
        Tuple of hidden states from all encoder layers.
    attentions (`tuple[torch.FloatTensor]`, *optional*):
        Tuple of attention weights from all encoder layers.
    Nr8   Úpos_embeds_flattenedÚtext_featuresÚspatial_shapesÚhidden_statesÚ
attentions)r+   r,   r-   r.   r8   r0   r1   r2   r=   r>   r?   Ú
LongTensorr@   r/   rA   r3   r4   r5   r<   r<   `   sµ   € € € € € € ðð ð ,0Ð�uÔ(Ð/Ð/Ñ/Ø59Ð˜%Ô+¨dÑ2Ð9Ð9Ñ9Ø.2€M�5Ô$ tÑ+Ð2Ð2Ñ2Ø.2€N�EÔ$ tÑ+Ð2Ð2Ñ2Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ø26€J��eÔ'Ô(¨4Ñ/Ð6Ð6Ñ6Ð6Ð6r4   r<   c                   ó°   — e Zd ZU dZdZej        ed<   dZej        ed<   dZ	ej        ed<   dZ
eej                 dz  ed<   dZeej                 dz  ed<   dS )ÚSam3DETRDecoderOutputa  
    intermediate_hidden_states (`torch.FloatTensor` of shape `(num_layers, batch_size, num_queries, hidden_size)`):
        Decoder hidden states from all layers.
    reference_boxes (`torch.FloatTensor` of shape `(num_layers, batch_size, num_queries, 4)`):
        Predicted reference boxes from all decoder layers in (cx, cy, w, h) format.
    presence_logits (`torch.FloatTensor` of shape `(num_layers, batch_size, 1)`):
        Presence logits from all decoder layers indicating object presence confidence.
    hidden_states (`tuple[torch.FloatTensor]`, *optional*):
        Tuple of hidden states from all decoder layers.
    attentions (`tuple[torch.FloatTensor]`, *optional*):
        Tuple of attention weights from all decoder layers (self-attention and cross-attention).
    NÚintermediate_hidden_statesÚreference_boxesÚpresence_logitsr@   rA   )r+   r,   r-   r.   rE   r0   r1   r2   rF   rG   r@   r/   rA   r3   r4   r5   rD   rD   z   s’   € € € € € € ðð ð 59Ð Ô 1Ð8Ð8Ñ8Ø)-€O�UÔ&Ð-Ð-Ñ-Ø)-€O�UÔ&Ð-Ð-Ñ-Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ø26€J��eÔ'Ô(¨4Ñ/Ð6Ð6Ñ6Ð6Ð6r4   rD   c                   ót   — e Zd ZU dZdZej        ed<   dZej        dz  ed<   dZ	e
ej                 dz  ed<   dS )ÚSam3MaskDecoderOutputaž  
    pred_masks (`torch.FloatTensor` of shape `(batch_size, num_queries, height, width)`):
        Predicted segmentation masks for each query.
    semantic_seg (`torch.FloatTensor` of shape `(batch_size, 1, height, width)`, *optional*):
        Semantic segmentation output.
    attentions (`tuple[torch.FloatTensor]`, *optional*):
        Tuple of attention weights from mask decoder cross-attention layers.
    NÚ
pred_masksÚsemantic_segrA   )r+   r,   r-   r.   rJ   r0   r1   r2   rK   rA   r/   r3   r4   r5   rI   rI   ‘   sf   € € € € € € ðð ð %)€J�Ô!Ð(Ð(Ñ(Ø-1€L�%Ô# dÑ*Ð1Ð1Ñ1Ø26€J��eÔ'Ô(¨4Ñ/Ð6Ð6Ñ6Ð6Ð6r4   rI   c                   óâ  — e Zd ZU dZdZej        ed<   dZej        ed<   dZ	ej        dz  ed<   dZ
ej        dz  ed<   dZej        dz  ed<   dZeej                 dz  ed<   dZej        dz  ed	<   dZeej                 dz  ed
<   dZeej                 dz  ed<   dZeej                 dz  ed<   dZeej                 dz  ed<   dZeej                 dz  ed<   dZeej                 dz  ed<   dS )ÚSam3ImageSegmentationOutputa}  
    pred_masks (`torch.FloatTensor` of shape `(batch_size, num_queries, height, width)`):
        Predicted segmentation masks for each query.
    pred_boxes (`torch.FloatTensor` of shape `(batch_size, num_queries, 4)`):
        Predicted bounding boxes in (x1, y1, x2, y2) format.
    pred_logits (`torch.FloatTensor` of shape `(batch_size, num_queries)`, *optional*):
        Classification confidence scores for each query, computed via dot product between
        decoder query features and text features.
    presence_logits (`torch.FloatTensor` of shape `(batch_size, 1)`, *optional*):
        Presence logits from the DETR decoder presence token (last layer only). These indicate whether objects
        are present in the scene. Can be used to compute final scores by multiplying with pred_logits:
        `final_scores = pred_logits.sigmoid() * presence_logits.sigmoid()`.
    semantic_seg (`torch.FloatTensor` of shape `(batch_size, 1, height, width)`, *optional*):
        Semantic segmentation output.
    decoder_hidden_states (`tuple[torch.FloatTensor]`, *optional*):
        Tuple of hidden states from all DETR decoder layers. Each tensor has shape `(batch_size, num_queries, hidden_size)`.
    decoder_reference_boxes (`torch.FloatTensor` of shape `(num_layers, batch_size, num_queries, 4)`, *optional*):
        Reference boxes from all DETR decoder layers.
    encoder_hidden_states (`tuple[torch.FloatTensor]`, *optional*):
        Tuple of hidden states from all DETR encoder layers.
    vision_hidden_states (`tuple[torch.FloatTensor]`, *optional*):
        Tuple of hidden states from all vision encoder (ViT) layers.
    vision_attentions (`tuple[torch.FloatTensor]`, *optional*):
        Attention weights from vision encoder (ViT) layers.
    detr_encoder_attentions (`tuple[torch.FloatTensor]`, *optional*):
        Attention weights from DETR encoder layers.
    detr_decoder_attentions (`tuple[torch.FloatTensor]`, *optional*):
        Attention weights from DETR decoder layers (self-attention and cross-attention).
    mask_decoder_attentions (`tuple[torch.FloatTensor]`, *optional*):
        Attention weights from mask decoder layers.
    NrJ   Ú
pred_boxesÚpred_logitsrG   rK   Údecoder_hidden_statesÚdecoder_reference_boxesÚencoder_hidden_statesÚvision_hidden_statesÚvision_attentionsÚdetr_encoder_attentionsÚdetr_decoder_attentionsÚmask_decoder_attentions)r+   r,   r-   r.   rJ   r0   r1   r2   rN   rO   rG   rK   rP   r/   rQ   rR   rS   rT   rU   rV   rW   r3   r4   r5   rM   rM   ¢   sx  € € € € € € ðð ð@ %)€J�Ô!Ð(Ð(Ñ(Ø$(€J�Ô!Ð(Ð(Ñ(Ø,0€K�Ô" TÑ)Ð0Ð0Ñ0Ø04€O�UÔ&¨Ñ-Ð4Ð4Ñ4Ø-1€L�%Ô# dÑ*Ð1Ð1Ñ1Ø=AÐ˜5 Ô!2Ô3°dÑ:ÐAÐAÑAØ8<Ð˜UÔ.°Ñ5Ð<Ð<Ñ<Ø=AÐ˜5 Ô!2Ô3°dÑ:ÐAÐAÑAØ<@Ð˜% Ô 1Ô2°TÑ9Ð@Ð@Ñ@Ø9=Ð�u˜UÔ.Ô/°$Ñ6Ð=Ð=Ñ=Ø?CÐ˜U 5Ô#4Ô5¸Ñ<ÐCÐCÑCØ?CÐ˜U 5Ô#4Ô5¸Ñ<ÐCÐCÑCØ?CÐ˜U 5Ô#4Ô5¸Ñ<ÐCÐCÑCÐCÐCr4   rM   çü©ñÒMbP?ÚxÚepsÚreturnc                 ó¼   — |                       dd¬¦  «        } |                       |¬¦  «        }d| z
                        |¬¦  «        }t          j        ||z  ¦  «        S )z5The inverse function for sigmoid activation function.r   r   ©ÚminÚmax©r^   )Úclampr0   Úlog)rY   rZ   Úx1Úx2s       r5   Úinverse_sigmoidre   Ô   sU   € à	�Š�A˜1ˆÑÔ€AØ	
�Š�SˆÑ	Ô	€BØ
ˆa‰%�Š˜3ˆÑ	Ô	€BÝŒ9�R˜"‘WÑÔÐr4   FÚreturn_indexc                 óŒ  — | j         \  }}}|j         \  }}	}
||cxk    r3|                     d¦  «        cxk    r|                     d¦  «        k    sn J ‚||
k    sJ ‚||                     d¦  «        k    sJ ‚|	|                     d¦  «        k    sJ ‚|                     d¬¦  «        }|                     d¬¦  «        }||z   }||	z   }t          j        ||j        ¬¦  «        d                              |d¦  «        |dd…df         k     }t          j        |||f|j        |j        ¬¦  «        }| |dd…d|…dd…f<   t          j        |	|j        ¬¦  «        d                              |d¦  «        }||dd…df         z   }| 	                    d|dd…dd…df          
                    dd|¦  «        |¦  «        }|r|||fS ||fS )aÇ  
    Concatenates two right-padded sequences, such that the resulting sequence
    is contiguous and also right-padded.

    Tensors are batch-first, masks are batch-first with True=valid, False=padding.

    Args:
        seq1: A tensor of shape (batch_size, seq1_length, hidden_size).
        mask1: A tensor of shape (batch_size, seq1_length) with True=valid, False=padding.
        seq2: A tensor of shape (batch_size, seq2_length, hidden_size).
        mask2: A tensor of shape (batch_size, seq2_length) with True=valid, False=padding.
        return_index: If True, also returns the index of the ids of the element of seq2
            in the concatenated sequence. This can be used to retrieve the elements of seq2.

    Returns:
        A tuple (concatenated_sequence, concatenated_mask) if return_index is False,
        otherwise (concatenated_sequence, concatenated_mask, index).
        The concatenated_mask uses True=valid, False=padding convention.
    r   r   éÿÿÿÿ©Údim)ÚdeviceN©rk   Údtype)ÚshapeÚsizeÚsumr0   Úarangerk   ÚrepeatÚzerosrm   ÚscatterÚexpand)Úseq1Úmask1Úseq2Úmask2rf   Ú
batch_sizeÚseq1_lengthÚhidden_sizeÚbatch_size2Úseq2_lengthÚhidden_size2Úactual_seq1_lengthsÚactual_seq2_lengthsÚfinal_lengthsÚ
max_lengthÚconcatenated_maskÚconcatenated_sequenceÚindexs                     r5   Úconcat_padded_sequencesr‡   Ü   s$  € ð( ,0¬:Ñ(€J�˜[Ø-1¬ZÑ*€K�˜là˜ÐFÐFÒFÐF¨¯
ª
°1©¬ÐFÐFÒFÐF¸¿ºÀA¹¼ÒFÐFÐFÐFÐFÐFØ˜,Ò&Ð&Ð&Ð&Ø˜%Ÿ*š* Q™-œ-Ò'Ð'Ð'Ð'Ø˜%Ÿ*š* Q™-œ-Ò'Ð'Ð'Ð'àŸ)š)¨˜)Ñ+Ô+ÐØŸ)š)¨˜)Ñ+Ô+Ðà'Ð*=Ñ=€MØ˜{Ñ*€Jõ 	Œ�Z¨¬Ð4Ñ4Ô4°TÔ:×AÒAÀ*ÈaÑPÔPÐS`ÐabÐabÐabÐdhÐahÔSiÒið õ "œK¨°ZÀÐ(MÐVZÔVaÐimÔisÐtÑtÔtÐØ04Ð˜!˜!˜!˜\˜k˜\¨1¨1¨1Ð,Ñ-õ ŒL˜¨T¬[Ð9Ñ9Ô9¸$Ô?×FÒFÀzÐSTÑUÔU€EØÐ'¨¨¨¨4¨Ô0Ñ0€Eð 2×9Ò9¸!¸UÀ1À1À1ÀaÀaÀaÈÀ:Ô=N×=UÒ=UÐVXÐZ\Ð^iÑ=jÔ=jÐlpÑqÔqÐàð ?Ø$Ð&7¸Ð>Ð>à Ð"3Ð3Ð3r4   c                 óž   — |                       d¦  «        \  }}}}|d|z  z
  |d|z  z
  |d|z  z   |d|z  z   g}t          j        |d¬¦  «        S )zDConvert boxes from (cx, cy, w, h) format to (x1, y1, x2, y2) format.rh   ç      à?ri   )Úunbindr0   Ústack)rY   Úx_cÚy_cÚwÚhÚbs         r5   Úbox_cxcywh_to_xyxyr‘     s\   € à—X’X˜b‘\”\�N€Cˆˆa�Ø
��a‘‰-˜3  q¡™=¨C°#¸±'©M¸SÀ3ÈÁ7¹]ÐL€AÝŒ;�q˜bÐ!Ñ!Ô!Ð!r4   c                   óB   ‡ — e Zd Zˆ fd„Zdej        dej        fd„Zˆ xZS )ÚSam3MLPc                 óP  •— t          ¦   «                              ¦   «          || _        t          |j                 | _        t          j        |j        |j	        ¦  «        | _
        t          j        |j	        |j        ¦  «        | _        t          j        |j        ¦  «        | _        d S ©N)ÚsuperÚ__init__Úconfigr   Ú
hidden_actÚactivation_fnÚnnÚLinearr|   Úintermediate_sizeÚfc1Úfc2ÚDropoutÚhidden_dropoutÚdropout©Úselfr˜   Ú	__class__s     €r5   r—   zSam3MLP.__init__  sz   ø€ Ý‰Œ×ÒÑÔÐØˆŒÝ# FÔ$5Ô6ˆÔÝ”9˜VÔ/°Ô1IÑJÔJˆŒÝ”9˜VÔ5°vÔ7IÑJÔJˆŒÝ”z &Ô"7Ñ8Ô8ˆŒˆˆr4   r@   r[   c                 ó®   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|S r•   )rž   r¢   rš   rŸ   )r¤   r@   s     r5   ÚforwardzSam3MLP.forward"  sN   € ØŸš Ñ/Ô/ˆØŸš ]Ñ3Ô3ˆØ×*Ò*¨=Ñ9Ô9ˆØŸš Ñ/Ô/ˆØÐr4   ©r+   r,   r-   r—   r0   r   r§   Ú__classcell__©r¥   s   @r5   r“   r“     s^   ø€ € € € € ð9ð 9ð 9ð 9ð 9ð U¤\ð °e´lð ð ð ð ð ð ð ð r4   r“   ç        ÚmoduleÚqueryÚkeyÚvaluer9   Úscalingr¢   Úkwargsc                 ó®  — |€|                      d¦  «        dz  }t          j        ||                     dd¦  «        ¦  «        |z  }|�||z   }t          j                             |d¬¦  «        }t          j                             ||| j        ¬¦  «        }t          j        ||¦  «        }	|	                     dd¦  «         	                    ¦   «         }	|	|fS )Nrh   ç      à¿r   r   ri   )ÚpÚtrainingr   )
ro   r0   ÚmatmulÚ	transposer›   Ú
functionalÚsoftmaxr¢   rµ   Ú
contiguous)
r¬   r­   r®   r¯   r9   r°   r¢   r±   Úattn_weightsÚattn_outputs
             r5   Úeager_attention_forwardr½   *  sÈ   € ð €Ø—*’*˜R‘.”. DÑ(ˆõ ”<  s§}¢}°Q¸Ñ':Ô':Ñ;Ô;¸gÑE€LàÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2Ð(Ñ>Ô>€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€Lå”,˜|¨UÑ3Ô3€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r4   c                   ó¤   ‡ — e Zd ZdZˆ fd„Z	 ddej        dej        dej        dej        dz  dee         d	e	ej        ej        f         fd
„Z
ˆ xZS )ÚSam3Attentionz`
    Multi-head attention.
    Handles standard [batch_size, seq_len, hidden_size] tensors.
    c                 óú  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        | j        |j        z  | _        | j        dz  | _        d| _        t          j	        | j        | j        ¦  «        | _
        t          j	        | j        | j        ¦  «        | _        t          j	        | j        | j        ¦  «        | _        t          j	        | j        | j        ¦  «        | _        d S ©Nr³   F)r–   r—   r˜   r|   Únum_attention_headsÚhead_dimr°   Ú	is_causalr›   rœ   Úq_projÚk_projÚv_projÚo_projr£   s     €r5   r—   zSam3Attention.__init__L  sÅ   ø€ Ý‰Œ×ÒÑÔÐØˆŒØ!Ô-ˆÔØ#)Ô#=ˆÔ ØÔ(¨FÔ,FÑFˆŒØ”} dÑ*ˆŒØˆŒå”i Ô 0°$Ô2BÑCÔCˆŒÝ”i Ô 0°$Ô2BÑCÔCˆŒÝ”i Ô 0°$Ô2BÑCÔCˆŒÝ”i Ô 0°$Ô2BÑCÔCˆŒˆˆr4   Nr­   r®   r¯   r9   r±   r[   c                 óÌ  — |j         d         }|j         d         }|j         d         }|                      |¦  «                             ||| j        | j        ¦  «                             dd¦  «        }|                      |¦  «                             ||| j        | j        ¦  «                             dd¦  «        }|                      |¦  «                             ||| j        | j        ¦  «                             dd¦  «        }t          j	        | j
        j        t          ¦  «        }	t          | j
        ¦  «        r>|�<|j        t          j        k    r't          d         }	t"                               d¦  «          |	| |||f|d| j        | j        dœ|¤Ž\  }
}|
                     ||| j        | j        z  ¦  «                             ¦   «         }
|                      |
¦  «        }
|
|fS )	aã  
        Args:
            query: [batch_size, query_len, hidden_size]
            key: [batch_size, key_len, hidden_size]
            value: [batch_size, value_len, hidden_size]
            attention_mask: [batch_size, num_heads, query_len, key_len] or broadcastable

        Returns:
            Tuple of (output, attention_weights)
                output: [batch_size, query_len, hidden_size]
                attention_weights: [batch_size, num_heads, query_len, key_len]
        r   r   r   NÚsdpaz‡Sam3Attention: falling back to SDPA for relative-position cross-attention because Flash Attention does not support additive bias masks.r«   ©r9   r¢   r°   rÄ   )rn   rÅ   ÚviewrÂ   rÃ   r·   rÆ   rÇ   r   Úget_interfacer˜   Ú_attn_implementationr½   r   rm   r0   ÚboolÚloggerÚwarning_oncer°   rÄ   Úreshaperº   rÈ   )r¤   r­   r®   r¯   r9   r±   rz   Ú	query_lenÚkey_lenÚattention_interfacer¼   r»   s               r5   r§   zSam3Attention.forwardZ  sê  € ð( ”[ ”^ˆ
Ø”K ”Nˆ	Ø”)˜A”,ˆà—’˜EÑ"Ô"×'Ò'¨
°I¸tÔ?WÐY]ÔYfÑgÔg×qÒqÐrsÐuvÑwÔwˆØ�kŠk˜#ÑÔ×#Ò# J°¸Ô9QÐSWÔS`ÑaÔa×kÒkÐlmÐopÑqÔqˆØ—’˜EÑ"Ô"×'Ò'¨
°G¸TÔ=UÐW[ÔWdÑeÔe×oÒoÐpqÐstÑuÔuˆå(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðõ
 )¨¬Ñ5Ô5ð	àÐ*ØÔ$­¬
Ò2Ð2õ #:¸&Ô"AÐÝ×ÒðHñô ð ð
 %8Ð$7ØØØØð	
%
ð
 *ØØ”LØ”nð
%
ð 
%
ð ð
%
ð 
%
Ñ!ˆ�\ð "×)Ò)¨*°iÀÔAYÐ\`Ô\iÑAiÑjÔj×uÒuÑwÔwˆØ—k’k +Ñ.Ô.ˆà˜LÐ(Ð(r4   r•   )r+   r,   r-   r.   r—   r0   r   r   r   r/   r§   r©   rª   s   @r5   r¿   r¿   F  sÀ   ø€ € € € € ðð ð
Dð Dð Dð Dð Dð& /3ð<)ð <)àŒ|ð<)ð Œ\ð<)ð Œ|ð	<)ð
 œ tÑ+ð<)ð Ð+Ô,ð<)ð 
ˆuŒ|˜Uœ\Ð)Ô	*ð<)ð <)ð <)ð <)ð <)ð <)ð <)ð <)r4   r¿   c            	       ó�   ‡ — e Zd ZdZd
dedededefˆ fd„Z ej	        ¦   «         de
ej        ej        f         fd	„¦   «         Zˆ xZS )ÚSam3ViTRotaryEmbeddingz 
    Vision Rotary Position Embedding for SAM3, following transformers library standards.
    Supports 2D (axial) rotary embeddings for spatial dimensions.
    ç      ð?r˜   Úend_xÚend_yÚscalec                 ó–  •— t          ¦   «                              ¦   «          |j        |j        z  }|dz  dk    rt	          d¦  «        ‚||c| _        | _        || _        |j        | _        || _	        d|j        t          j        d|d¦  «        d |dz  …                              ¦   «         |z  z  z  }t          j        ||z  t          j        ¬¦  «        }||z  |z  }t          j        ||d¬¦  «        |z  }	t          j        ||¦  «                             ¦   «         }
t          j        |	|¦  «                             ¦   «         }t          j        |
|gd¬	¦  «        }|                     d
d¬	¦  «        }|                      d|                     ¦   «         d¬¦  «         |                      d|                     ¦   «         d¬¦  «         d S )Né   r   z/Dimension must be divisible by 4 for axial RoPErØ   ©rm   Úfloor©Úrounding_moderh   ri   r   Úrope_embeddings_cosF)Ú
persistentÚrope_embeddings_sin)r–   r—   r|   rÂ   Ú
ValueErrorrÙ   rÚ   rj   Ú
rope_thetarÛ   r0   rq   ÚfloatÚlongÚdivÚouterÚcatÚrepeat_interleaveÚregister_bufferÚcosÚsin)r¤   r˜   rÙ   rÚ   rÛ   rj   ÚfreqsÚflattened_indicesÚx_positionsÚy_positionsÚfreqs_xÚfreqs_yÚinv_freqr¥   s                €r5   r—   zSam3ViTRotaryEmbedding.__init__Ÿ  s¬  ø€ Ý‰Œ×ÒÑÔÐØÔ  FÔ$>Ñ>ˆà�‰7�aŠ<ˆ<ÝÐNÑOÔOÐOØ!&¨ÐˆŒ
�D”JØˆŒØ Ô+ˆŒØˆŒ
Ø�vÔ(­U¬\¸!¸SÀ!Ñ-DÔ-DÀ\ÈÈqÉÀ\Ô-R×-XÒ-XÑ-ZÔ-ZÐ]`Ñ-`ÑaÑbˆå!œL¨°©½e¼jÐIÑIÔIÐØ(¨5Ñ0°EÑ9ˆÝ”iÐ 1°5ÈÐPÑPÔPÐSXÑXˆÝ”+˜k¨5Ñ1Ô1×7Ò7Ñ9Ô9ˆÝ”+˜k¨5Ñ1Ô1×7Ò7Ñ9Ô9ˆÝ”9˜g wÐ/°RÐ8Ñ8Ô8ˆØ×-Ò-¨a°RÐ-Ñ8Ô8ˆà×ÒÐ2°H·L²L±N´NÈuÐÑUÔUÐUØ×ÒÐ2°H·L²L±N´NÈuÐÑUÔUÐUÐUÐUr4   r[   c                 ó   — | j         | j        fS r•   )râ   rä   ©r¤   s    r5   r§   zSam3ViTRotaryEmbedding.forward¶  s   € ð Ô'¨Ô)AÐAÐAr4   )rØ   )r+   r,   r-   r.   r&   Úintrç   r—   r0   Úno_gradr/   r   r§   r©   rª   s   @r5   r×   r×   ™  s¹   ø€ € € € € ðð ð
Vð V˜}ð V°Sð VÀð VÈUð Vð Vð Vð Vð Vð Vð. €U„]�_„_ðB˜˜uœ|¨U¬\Ð9Ô:ð Bð Bð Bñ „_ðBð Bð Bð Bð Br4   r×   c                 óÎ   —  | j         g | j        dd…         ¢d‘d‘R Ž } |                      d¬¦  «        \  }}t          j        | |fd¬¦  «        } |                      d¬¦  «        S )ax  
    pairwise rotation of the hidden dims of the input. Differerent from Llama Half-Tensor Rotation.

    This is an optimized version of the following more explicit implementation:
    ```python
    x_rotated = torch.zeros_like(x, dtype=x.dtype, device=x.device)
    x_rotated[..., ::2] = -x[..., 1::2]
    x_rotated[..., 1::2] = x[..., ::2]
    return x_rotated
    ```
    Nrh   r   ri   éþÿÿÿ)Ú	start_dim)rÌ   rn   rŠ   r0   r‹   Úflatten)rY   rc   rd   s      r5   Úrotate_pairwiserÿ   ¼  st   € ð 	ˆŒÐ$�”˜˜˜”Ð$˜bÐ$ !Ð$Ð$Ð$€AØ�XŠX˜"ˆXÑÔ�F€BˆÝŒ�b�S˜"�I 2Ð&Ñ&Ô&€AØ�9Š9˜rˆ9Ñ"Ô"Ð"r4   ÚqÚkrî   rï   c                 ó  — |                       ¦   «         }||z  t          |¦  «        |z  z   }|                      ¦   «         }||z  t          |¦  «        |z  z   }|                     | ¦  «        |                     |¦  «        fS )aÄ  
    Apply rotary position embedding to query and key tensors for self-attention.

    Args:
        q: Query tensor of shape (batch_size, num_windows, seq_len, num_heads, head_dim)
        k: Key tensor of shape (batch_size, num_windows, seq_len, num_heads, head_dim)
        cos: Cosine position embedding of shape (seq_len, head_dim)
        sin: Sine position embedding of shape (seq_len, head_dim)

    Returns:
        Rotated (q, k) tensors
    )rç   rÿ   Útype_as)r   r  rî   rï   Úq_embedÚk_embeds         r5   Úapply_rotary_pos_emb_2dr  Î  sw   € ð$ �gŠg‰iŒi€GØ˜‰}¥°Ñ!9Ô!9¸CÑ!?Ñ@€Gà�gŠg‰iŒi€GØ˜‰}¥°Ñ!9Ô!9¸CÑ!?Ñ@€Gà�?Š?˜1ÑÔ˜wŸš¨qÑ1Ô1Ð1Ð1r4   c                   óz   ‡ — e Zd ZdZdefˆ fd„Zdej        deej        ej        f         de	e
         defd„Zˆ xZS )	ÚSam3ViTRoPEAttentionz-Self-attention with rotary position encoding.r˜   c                 ó  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        | j        |j        z  | _        | j        dz  | _        |j        | _        d| _        t          j
        | j        | j        ¦  «        | _        t          j
        | j        | j        ¦  «        | _        t          j
        | j        | j        ¦  «        | _        t          j
        | j        | j        ¦  «        | _        d S rÁ   )r–   r—   r˜   r|   rÂ   rÃ   r°   Úattention_dropoutrÄ   r›   rœ   rÅ   rÆ   rÇ   rÈ   r£   s     €r5   r—   zSam3ViTRoPEAttention.__init__ì  sÐ   ø€ Ý‰Œ×ÒÑÔÐØˆŒØ!Ô-ˆÔØ#)Ô#=ˆÔ ØÔ(¨FÔ,FÑFˆŒØ”} dÑ*ˆŒØ!'Ô!9ˆÔØˆŒå”i Ô 0°$Ô2BÑCÔCˆŒÝ”i Ô 0°$Ô2BÑCÔCˆŒÝ”i Ô 0°$Ô2BÑCÔCˆŒÝ”i Ô 0°$Ô2BÑCÔCˆŒˆˆr4   r@   Úposition_embeddingsr±   r[   c                 óÆ  — |j         \  }}}}||z  }||| j        | j        f}	 |                      |¦  «        j        |	Ž                      dd¦  «        }
 |                      |¦  «        j        |	Ž                      dd¦  «        } |                      |¦  «        j        |	Ž                      dd¦  «        }|\  }}t          |
|||¬¦  «        \  }
}t          j
        | j        j        t          ¦  «        } || |
||fd | j        sdn| j        | j        | j        dœ|¤Ž\  }}|                     |||d¦  «                             ¦   «         }|                      |¦  «        }||fS )Nr   r   )rî   rï   r«   rË   rh   )rn   rÂ   rÃ   rÅ   rÌ   r·   rÆ   rÇ   r  r   rÍ   r˜   rÎ   r½   rµ   r
  r°   rÄ   rÒ   rº   rÈ   )r¤   r@   r  r±   rz   ÚheightÚwidthÚ_Úseq_lenÚ	new_shaper­   r®   r¯   rî   rï   rÕ   r¼   r»   s                     r5   r§   zSam3ViTRoPEAttention.forwardû  s�  € ð (5Ô':Ñ$ˆ
�F˜E 1Ø˜5‘.ˆØ ¨$Ô*BÀDÄMÐRˆ	Ø/�—’˜MÑ*Ô*Ô/°Ð;×EÒEÀaÈÑKÔKˆØ-ˆd�kŠk˜-Ñ(Ô(Ô-¨yÐ9×CÒCÀAÀqÑIÔIˆØ/�—’˜MÑ*Ô*Ô/°Ð;×EÒEÀaÈÑKÔKˆØ&‰ˆˆSÝ,¨U°C¸SÀcÐJÑJÔJ‰
ˆˆså(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð %8Ð$7ØØØØð	
%
ð
  Ø#œ}ÐH�C�C°$Ô2HØ”LØ”nð
%
ð 
%
ð ð
%
ð 
%
Ñ!ˆ�\ð "×)Ò)¨*°f¸eÀRÑHÔH×SÒSÑUÔUˆØ—k’k +Ñ.Ô.ˆØ˜LÐ(Ð(r4   )r+   r,   r-   r.   r&   r—   r0   r   r/   r   r   r§   r©   rª   s   @r5   r  r  é  s¡   ø€ € € € € Ø7Ð7ðD˜}ð Dð Dð Dð Dð Dð Dð )à”|ð )ð # 5¤<°´Ð#=Ô>ð )ð Ð+Ô,ð	 )ð
 
ð )ð  )ð  )ð  )ð  )ð  )ð  )ð  )r4   r  c                   óL   ‡ — e Zd ZdZdefˆ fd„Zdej        dej        fd„Zˆ xZ	S )ÚSam3ViTPatchEmbeddingszì
    This class turns `pixel_values` of shape `(batch_size, num_channels, height, width)` into the initial
    `hidden_states` (patch embeddings) of shape `(batch_size, seq_length, hidden_size)` to be consumed by a
    Transformer.
    r˜   c                 ó¦  •— t          ¦   «                              ¦   «          |j        |j        }}|j        |j        }}t          |t          ¦  «        r|n||f}t          |t          ¦  «        r|n||f}|d         |d         z  |d         |d         z  z  }|| _        || _        || _        || _	        t          j        ||||d¬¦  «        | _        d S )Nr   r   F)Úkernel_sizeÚstrideÚbias)r–   r—   Úpretrain_image_sizeÚ
patch_sizeÚnum_channelsr|   Ú
isinstancer   Ú
image_sizeÚnum_patchesr›   ÚConv2dÚ
projection)r¤   r˜   r  r  r  r|   r  r¥   s          €r5   r—   zSam3ViTPatchEmbeddings.__init__%  sØ   ø€ Ý‰Œ×ÒÑÔÐØ!'Ô!;¸VÔ=N�Jˆ
Ø$*Ô$7¸Ô9K�kˆå#-¨j½(Ñ#CÔ#CÐa�Z�ZÈ*ÐV`ÐIaˆ
Ý#-¨j½(Ñ#CÔ#CÐa�Z�ZÈ*ÐV`ÐIaˆ
Ø! !”}¨
°1¬Ñ5¸*ÀQ¼-È:ÐVWÌ=Ñ:XÑYˆØ$ˆŒØ$ˆŒØ(ˆÔØ&ˆÔåœ) L°+È:Ð^hÐotÐuÑuÔuˆŒˆˆr4   Úpixel_valuesr[   c                 óÂ   — |                       |                     | j         j        j        ¦  «        ¦  «                             d¦  «                             dd¦  «        }|S )Nr   r   )r  ÚtoÚweightrm   rþ   r·   )r¤   r   Ú
embeddingss      r5   r§   zSam3ViTPatchEmbeddings.forward4  sN   € Ø—_’_ \§_¢_°T´_Ô5KÔ5QÑ%RÔ%RÑSÔS×[Ò[Ð\]Ñ^Ô^×hÒhÐijÐlmÑnÔnˆ
ØÐr4   )
r+   r,   r-   r.   r&   r—   r0   r   r§   r©   rª   s   @r5   r  r    s{   ø€ € € € € ðð ðv˜}ð vð vð vð vð vð vð E¤Lð °U´\ð ð ð ð ð ð ð ð r4   r  c                   ó€   ‡ — e Zd ZdZdefˆ fd„Zdej        dededej        fd„Z		 dd
ej        de
dej        fd„Zˆ xZS )ÚSam3ViTEmbeddingsz²
    Construct the patch embeddings and position embeddings for SAM3 ViT.

    Position embeddings are tiled (not interpolated) when resizing to match different input sizes.
    r˜   c                 ó@  •— t          ¦   «                              ¦   «          t          |¦  «        | _        | j        j        }t          j        t          j        d||j	        ¦  «        ¦  «        | _
        t          j        |j        ¦  «        | _        |j        | _        d S )Nr   )r–   r—   r  Úpatch_embeddingsr  r›   Ú	Parameterr0   Úrandnr|   r  r    r¡   r¢   r  )r¤   r˜   r  r¥   s      €r5   r—   zSam3ViTEmbeddings.__init__@  s€   ø€ Ý‰Œ×ÒÑÔÐå 6°vÑ >Ô >ˆÔØÔ+Ô7ˆÝ#%¤<ÝŒK˜˜;¨Ô(:Ñ;Ô;ñ$
ô $
ˆÔ õ ”z &Ô"7Ñ8Ô8ˆŒØ Ô+ˆŒˆˆr4   r  r  r  r[   c                 ó  — t          |j        d         dz  ¦  «        }t          j                             ¦   «         s&||k    r ||k    r|                     d||z  d¦  «        S |j        d         }|                     d|||¦  «                             dddd¦  «        }||z  dz   }||z  dz   }|                     dd||g¦  «        dd…dd…d|…d|…f         }|                     dddd¦  «                             d||z  |¦  «        S )aG  
        Tile position embeddings to match target spatial dimensions.
        Args:
            position_embeddings: Shape [1, num_pretrain_patches, hidden_size]
            height: Target height in patches
            width: Target width in patches

        Returns:
            Shape [1, height * width, hidden_size]
        r   r‰   rh   r   r   r   N)rù   rn   r0   ÚjitÚ
is_tracingrÒ   ÚpermuteÚtile)	r¤   r  r  r  Úpretrain_sizer|   Ú	pos_embedÚrepeat_hÚrepeat_ws	            r5   Ú_tile_position_embeddingsz+Sam3ViTEmbeddings._tile_position_embeddingsL  s0  € õ  Ð/Ô5°aÔ8¸CÑ?Ñ@Ô@ˆõ Œy×#Ò#Ñ%Ô%ð 	F¨-¸6Ò*AÐ*AÀmÐW\ÒF\ÐF\Ø&×.Ò.¨q°&¸5±.À"ÑEÔEÐEð *Ô/°Ô3ˆØ'×/Ò/°°=À-ÐQ\Ñ]Ô]×eÒeÐfgÐijÐlmÐopÑqÔqˆ	Ø˜]Ñ*¨QÑ.ˆØ˜MÑ)¨AÑ-ˆØ—N’N A q¨(°HÐ#=Ñ>Ô>¸q¸q¸qÀ!À!À!ÀWÀfÀWÈfÈuÈfÐ?TÔUˆ	Ø× Ò   A q¨!Ñ,Ô,×4Ò4°Q¸À¹ÈÑTÔTÐTr4   Fr   Úinterpolate_pos_encodingc                 óè   — |j         dd …         \  }}|                      |¦  «        }|| j        z  }|| j        z  }|                      | j        ||¦  «        }||z   }|                      |¦  «        }|S )Nrü   )rn   r(  r  r4  r  r¢   )	r¤   r   r5  r  r  r$  Úheight_patchesÚwidth_patchesr  s	            r5   r§   zSam3ViTEmbeddings.forwardj  sŒ   € ð
 %Ô*¨2¨3¨3Ô/‰ˆ�Ø×*Ò*¨<Ñ8Ô8ˆ
ð   4¤?Ñ2ˆØ ¤Ñ0ˆà"×<Ò<ØÔ$ØØñ
ô 
Ðð
  Ð"5Ñ5ˆ
Ø—\’\ *Ñ-Ô-ˆ
àÐr4   ©F)r+   r,   r-   r.   r&   r—   r0   r   rù   r4  rÏ   r§   r©   rª   s   @r5   r&  r&  9  sÓ   ø€ € € € € ðð ð
,˜}ð 
,ð 
,ð 
,ð 
,ð 
,ð 
,ðUà"œ\ðUð ðUð ð	Uð
 
ŒðUð Uð Uð UðB */ðð à”lðð #'ðð 
Œð	ð ð ð ð ð ð ð r4   r&  c           	      óv  — | j         \  }}}}|||z  z
  |z  }|||z  z
  |z  }t          j                             | ddd|d|f¦  «        } ||z   ||z   }	}|                      |||z  ||	|z  ||¦  «        } |                      dddddd¦  «                             ¦   «                              d|||¦  «        }
|
||	ffS )a  
    Partition into non-overlapping windows with padding if needed.

    Args:
        hidden_state (`torch.Tensor`):
            Input tokens with [batch_size, height, width, num_channels].
        window_size (`int`):
            Window size.

    Returns:
        `tuple(torch.FloatTensor)` comprising various elements:
        - windows: windows after partition with [batch_size * num_windows, window_size, window_size, num_channels].
        - (padded_height, padded_width): padded height and width before partition
    r   r   r   r   rÝ   é   rh   )rn   r›   r¸   ÚpadrÌ   r.  rº   )Úhidden_stateÚwindow_sizerz   r  r  r  Ú
pad_heightÚ	pad_widthÚpadded_heightÚpadded_widthÚwindowss              r5   Úwindow_partitionrD  �  sõ   € ð /;Ô.@Ñ+€J�˜˜|Ø ¨Ñ 4Ñ4¸ÑC€JØ˜u {Ñ2Ñ2°kÑA€Iõ ”=×$Ò$ \°A°q¸!¸YÈÈ:Ð3VÑWÔW€Là"(¨:Ñ"5°u¸yÑ7H�<€Mà×$Ò$Ø�M [Ñ0°+¸|È{Ñ?ZÐ\gÐiuñô €Lð ×"Ò" 1 a¨¨A¨q°!Ñ4Ô4×?Ò?ÑAÔA×FÒFÀrÈ;ÐXcÐeqÑrÔr€GØ�] LÐ1Ð1Ð1r4   c                 ó`  — |\  }}|\  }}| j         d         ||z  |z  |z  z  }|                      |||z  ||z  ||d¦  «        }	|	                     dddddd¦  «                             ¦   «         }	|	                     |||d¦  «        }	|	dd…d|…d|…dd…f                              ¦   «         }	|	S )	aB  
    Window unpartition into original sequences and removing padding.

    Args:
        windows (`torch.Tensor`):
            Input tokens with [batch_size * num_windows, window_size, window_size, num_channels].
        window_size (`int`):
            Window size.
        pad_height_width (`tuple[int]`):
            Padded height and width (padded_height, padded_width).
        height_width (`tuple[int]`):
            Original height and width before padding.

    Returns:
        hidden_state: unpartitioned sequences with [batch_size, height, width, num_channels].
    r   rh   r   r   r   rÝ   r;  N)rn   rÌ   r.  rº   )
rC  r>  Úpad_height_widthÚheight_widthrA  rB  r  r  rz   r=  s
             r5   Úwindow_unpartitionrH     sß   € ð" #3Ñ€M�<Ø �M€FˆEØ”˜qÔ! m°lÑ&BÀkÑ&QÐU`Ñ&`Ña€JØ—<’<Ø�M [Ñ0°,À+Ñ2MÈ{Ð\gÐikñô €Lð  ×'Ò'¨¨1¨a°°A°qÑ9Ô9×DÒDÑFÔF€LØ×$Ò$ Z°ÀÈbÑQÔQ€Lð      7 F 7¨F¨U¨F°A°A°AÐ 5Ô6×AÒAÑCÔC€LØÐr4   c                   óD   ‡ — e Zd Zdˆ fd„Zdej        dej        fd„Zˆ xZS )ÚSam3ViTLayerScaler[   Nc                 ó¸   •— t          ¦   «                              ¦   «          t          j        |j        t          j        |j        ¦  «        z  ¦  «        | _        d S r•   )	r–   r—   r›   r)  Úlayer_scale_init_valuer0   Úonesr|   Úlambda1r£   s     €r5   r—   zSam3ViTLayerScale.__init__À  sC   ø€ Ý‰Œ×ÒÑÔÐÝ”| FÔ$AÅEÄJÈvÔOaÑDbÔDbÑ$bÑcÔcˆŒˆˆr4   r=  c                 ó   — || j         z  S r•   )rN  )r¤   r=  s     r5   r§   zSam3ViTLayerScale.forwardÄ  s   € Ø˜dœlÑ*Ð*r4   )r[   Nr¨   rª   s   @r5   rJ  rJ  ¿  si   ø€ € € € € ðdð dð dð dð dð dð+ E¤Lð +°U´\ð +ð +ð +ð +ð +ð +ð +ð +r4   rJ  c                   óf   ‡ — e Zd ZdZddededdfˆ fd„Zdej        d	e	e
         dej        fd
„Zˆ xZS )ÚSam3ViTLayerzYVision Transformer layer with rotary position embeddings and optional windowed attention.r   r˜   r>  r[   Nc                 óØ  •— t          ¦   «                              ¦   «          |j        }|j        }t	          |t
          t          f¦  «        r|n||f}|j        }t	          |t
          t          f¦  «        r|n||f}|d         |d         z  |d         |d         z  f}t          j	        ||j
        ¬¦  «        | _        |dk    r|n||f}|j        |d         z  }t          ||d         |d         |¬¦  «        | _        t          |¦  «        | _        t          j	        ||j
        ¬¦  «        | _        t%          |¦  «        | _        t          j        |j        ¦  «        | _        || _        d S )Nr   r   ©rZ   )rÙ   rÚ   rÛ   )r–   r—   r|   r  r  Úlistr/   r  r›   Ú	LayerNormÚlayer_norm_epsÚlayer_norm1r>  r×   Ú
rotary_embr  Ú	attentionÚlayer_norm2r“   Úmlpr    r¡   r¢   )
r¤   r˜   r>  r|   r  r  Ú
input_sizeÚrotary_input_sizeÚrotary_scaler¥   s
            €r5   r—   zSam3ViTLayer.__init__Ë  sf  ø€ Ý‰Œ×ÒÑÔÐàÔ(ˆØÔ&ˆ
Ý#-¨j½4Å¸-Ñ#HÔ#HÐf�Z�ZÈzÐ[eÐNfˆ
àÔ&ˆ
Ý#-¨j½4Å¸-Ñ#HÔ#HÐf�Z�ZÈzÐ[eÐNfˆ
à  ”m z°!¤}Ñ4°jÀ´mÀzÐRSÄ}Ñ6TÐUˆ
Ýœ<¨¸Ô9NÐOÑOÔOˆÔØ*5¸Ò*:Ð*:˜J˜JÀÈkÐ@ZÐØÔ)Ð,=¸aÔ,@Ñ@ˆÝ0ØÐ+¨AÔ.Ð6GÈÔ6JÐR^ð
ñ 
ô 
ˆŒõ .¨fÑ5Ô5ˆŒÝœ<¨¸Ô9NÐOÑOÔOˆÔÝ˜6‘?”?ˆŒÝ”z &Ô"7Ñ8Ô8ˆŒà&ˆÔÐÐr4   r@   r±   c                 óÔ  — |}|                       |¦  «        }| j        dk    r2|j        d         |j        d         }}t          || j        ¦  «        \  }}|                      ¦   «         } | j        ||fi |¤Ž\  }}| j        dk    rt          || j        |||f¦  «        }||z   }|}|                      |¦  «        }|                      |¦  «        }||  	                    |¦  «        z   }|S )Nr   r   r   )
rW  r>  rn   rD  rX  rY  rH  rZ  r[  r¢   )	r¤   r@   r±   Úresidualr  r  rF  r  r  s	            r5   r§   zSam3ViTLayer.forwardã  s  € ð
 !ˆà×(Ò(¨Ñ7Ô7ˆàÔ˜aÒÐØ)Ô/°Ô2°MÔ4GÈÔ4J�EˆFå.>¸}ÈdÔN^Ñ._Ô._Ñ+ˆMÐ+à"ŸošoÑ/Ô/ÐØ)˜4œ>¨-Ð9LÐWÐWÐPVÐWÐWÑˆ�qàÔ˜aÒÐå.¨}¸dÔ>NÐP`ÐciÐkpÐbqÑrÔrˆMà  =Ñ0ˆØ ˆØ×(Ò(¨Ñ7Ô7ˆØŸš Ñ/Ô/ˆØ  4§<¢<°Ñ#>Ô#>Ñ>ˆàÐr4   )r   )r+   r,   r-   r.   r&   rù   r—   r0   r   r   r   r§   r©   rª   s   @r5   rQ  rQ  È  s—   ø€ € € € € ØcÐcð'ð '˜}ð '¸3ð 'Àtð 'ð 'ð 'ð 'ð 'ð 'ð0à”|ðð Ð+Ô,ðð 
Œð	ð ð ð ð ð ð ð r4   rQ  )r0   Útorchvision)Úbackendsc                   óB   ‡ — e Zd ZeZdZdZddgZdZdZ	dZ
dZˆ fd„Zˆ xZS )ÚSam3PreTrainedModelÚsam3r   ÚimageÚtextTc                 óè  •— t          ¦   «                              |¦  «         t          |t          ¦  «        r(t	          j        |j        d| j        j        ¬¦  «         d S t          |t          ¦  «        �r||j
        |j        }}|j        }d|j        t          j        d|d¦  «        d |dz  …                              ¦   «         |z  z  z  }t          j        ||z  t          j        ¬¦  «        }||z  |j        z  }t          j        ||d¬¦  «        |j        z  }t          j        ||¦  «                             ¦   «         }	t          j        ||¦  «                             ¦   «         }
t          j        |	|
gd	¬
¦  «        }|                     dd	¬
¦  «        }t	          j        |j        |                     ¦   «         ¦  «         t	          j        |j        |                     ¦   «         ¦  «         d S d S )Nr«   )ÚmeanÚstdrØ   r   rÝ   rÞ   rß   rà   rh   ri   r   )r–   Ú_init_weightsr  r&  ÚinitÚnormal_r  r˜   Úinitializer_ranger×   rÙ   rÚ   rj   ræ   r0   rq   rç   rè   rÛ   ré   rê   rë   rì   Úcopy_râ   rî   rä   rï   )r¤   r¬   rÙ   rÚ   rj   rð   rñ   rò   ró   rô   rõ   rö   r¥   s               €r5   rk  z!Sam3PreTrainedModel._init_weights  s¿  ø€ Ý‰Œ×Ò˜fÑ%Ô%Ð%Ý�fÕ/Ñ0Ô0ð 	CÝŒL˜Ô3¸#À4Ä;ÔC`ÐaÑaÔaÐaÐaÐaÝ˜Õ 6Ñ7Ô7ñ 	CØ!œ<¨¬�5ˆEØ”*ˆCØ˜6Ô,µ´¸aÀÀaÑ1HÔ1HÈÈCÐSTÉHÈÔ1V×1\Ò1\Ñ1^Ô1^ÐadÑ1dÑeÑfˆEÝ %¤¨U°U©]Å%Ä*Ð MÑ MÔ MÐØ,¨uÑ4¸¼ÑDˆKÝœ)Ð$5°uÈGÐTÑTÔTÐW]ÔWcÑcˆKÝ”k +¨uÑ5Ô5×;Ò;Ñ=Ô=ˆGÝ”k +¨uÑ5Ô5×;Ò;Ñ=Ô=ˆGÝ”y '¨7Ð!3¸Ð<Ñ<Ô<ˆHØ×1Ò1°!¸Ð1Ñ<Ô<ˆHÝŒJ�vÔ1°8·<²<±>´>ÑBÔBÐBÝŒJ�vÔ1°8·<²<±>´>ÑBÔBÐBÐBÐBð	Cð 	Cr4   )r+   r,   r-   r    Úconfig_classÚbase_model_prefixÚmain_input_nameÚinput_modalitiesÚ_supports_sdpaÚ_supports_flash_attnÚ_supports_flex_attnÚ_supports_attention_backendrk  r©   rª   s   @r5   rd  rd    su   ø€ € € € € ð €LØÐØ$€OØ Ð(ÐØ€NØÐØÐØ"&ÐðCð Cð Cð Cð Cð Cð Cð Cð Cr4   rd  c            	       ó´   ‡ — e Zd ZU eed<   eedœZdefˆ fd„Zde	fd„Z
e ed¬¦  «        edej        d	ee         defd
„¦   «         ¦   «         ¦   «         Zˆ xZS )ÚSam3ViTModelr˜   ©r@   rA   c                 ób  •‡— t          ¦   «                              ‰¦  «         ‰| _        t          ‰¦  «        | _        t          j        ‰j        ‰j        ¬¦  «        | _	        t          j
        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        |                      ¦   «          d S )NrS  c                 óR   •— g | ]#}t          ‰|‰j        vr‰j        nd ¬¦  «        ‘Œ$S )r   )r>  )rQ  Úglobal_attn_indexesr>  )Ú.0Úir˜   s     €r5   ú
<listcomp>z)Sam3ViTModel.__init__.<locals>.<listcomp>.  sM   ø€ ð ð ð àõ ˜VÀqÐPVÔPjÐGjÐGj°Ô1CÐ1CÐpqÐrÑrÔrðð ð r4   )r–   r—   r˜   r&  r$  r›   rU  r|   rV  Ú
layer_normÚ
ModuleListÚrangeÚnum_hidden_layersÚlayersÚ	post_initr£   s    `€r5   r—   zSam3ViTModel.__init__(  sª   øø€ Ý‰Œ×Ò˜Ñ Ô Ð ØˆŒÝ+¨FÑ3Ô3ˆŒÝœ, vÔ'9¸vÔ?TÐUÑUÔUˆŒÝ”mðð ð ð å˜vÔ7Ñ8Ô8ðñ ô ñ
ô 
ˆŒð 	�ŠÑÔÐÐÐr4   r[   c                 ó   — | j         j        S r•   )r$  r(  rø   s    r5   Úget_input_embeddingsz!Sam3ViTModel.get_input_embeddings5  s   € ØŒÔ/Ð/r4   F)Útie_last_hidden_statesr   r±   c                 óœ  — |                       |¦  «        }|j        d         }|j        d         | j        j        z  }|j        d         | j        j        z  }|j        d         }|                     ||||¦  «        }|                      |¦  «        }| j        D ]} ||fi |¤Ž}Œ|                     |||z  |¦  «        }t          |¬¦  «        S )Nr   rü   rh   )r8   )r$  rn   r˜   r  rÌ   r�  r…  r   )	r¤   r   r±   r@   rz   r  r  r|   Úlayers	            r5   r§   zSam3ViTModel.forward8  sá   € ð Ÿš¨Ñ5Ô5ˆà"Ô(¨Ô+ˆ
ØÔ# BÔ'¨4¬;Ô+AÑAˆØÔ" 2Ô&¨$¬+Ô*@Ñ@ˆØ#Ô)¨"Ô-ˆð &×*Ò*¨:°v¸uÀkÑRÔRˆàŸš¨Ñ6Ô6ˆØ”[ð 	;ð 	;ˆEØ!˜E -Ð:Ð:°6Ð:Ð:ˆMˆMð &×*Ò*¨:°vÀ±~À{ÑSÔSˆå°Ð?Ñ?Ô?Ð?r4   )r+   r,   r-   r&   r2   rQ  r  Ú_can_record_outputsr—   r  rˆ  r   r   r   r0   r   r   r   r   r§   r©   rª   s   @r5   ry  ry     só   ø€ € € € € € àÐÐÑà%Ø*ðð Ðð
˜}ð ð ð ð ð ð ð0Ð&<ð 0ð 0ð 0ð 0ð  Ø€_¨EÐ2Ñ2Ô2Øð@à”lð@ð Ð+Ô,ð@ð 
ð	@ð @ð @ñ „^ñ 3Ô2ñ  Ôð@ð @ð @ð @ð @r4   ry  c                   óÀ  ‡ — e Zd ZdZ	 	 	 	 ddededed	edz  fˆ fd
„Zdej	        dej	        de
ej	        ej	        f         fd„Zdej	        dej	        fd„Ze ed¬¦  «        	 	 	 	 ddej        dej        ez  dej        deded	edz  dedej	        dz  dej	        fd„¦   «         ¦   «         Z	 ddej        dej        ez  dej        dej	        dz  dej	        f
d„Zˆ xZS )ÚSam3SinePositionEmbeddingz¬
    This is a more standard version of the position embedding, very similar to the one used by the Attention is all you
    need paper, generalized to work on images.
    é@   é'  FNÚnum_position_featuresÚtemperatureÚ	normalizerÛ   c                 óÌ   •— t          ¦   «                              ¦   «          |�|du rt          d¦  «        ‚|| _        || _        || _        |€dt          j        z  n|| _        d S )NFz+normalize should be True if scale is passedr   )	r–   r—   rå   r‘  r’  r“  ÚmathÚpirÛ   )r¤   r‘  r’  r“  rÛ   r¥   s        €r5   r—   z"Sam3SinePositionEmbedding.__init__Z  sj   ø€ õ 	‰Œ×ÒÑÔÐØÐ ¨eÐ!3Ð!3ÝÐJÑKÔKÐKØ%:ˆÔ"Ø&ˆÔØ"ˆŒØ$) M�Q�œ‘[�[°uˆŒ
ˆ
ˆ
r4   rY   Úyr[   c                 óÈ  — || j         z  }|| j         z  }t          j        | j        t          j        |j        ¬¦  «                             |j        ¦  «        }| j        d|dz  z  | j        z  z  }|dd…df         |z  }|dd…df         |z  }t          j	        |dd…ddd…f          
                    ¦   «         |dd…ddd…f                              ¦   «         fd¬¦  «                             d¦  «        }t          j	        |dd…ddd…f          
                    ¦   «         |dd…ddd…f                              ¦   «         fd¬¦  «                             d¦  «        }||fS )a  
        Encode 1D coordinate pairs using sine/cosine positional embeddings.

        Args:
            x: 1D tensor of x coordinates (flattened)
            y: 1D tensor of y coordinates (flattened)

        Returns:
            Tuple of (pos_x, pos_y) positional embeddings
        ©rm   rk   r   Nr   r   ri   )rÛ   r0   rq   r‘  Úint64rk   r"  rm   r’  r‹   rï   rî   rþ   )r¤   rY   r—  Úx_embedÚy_embedÚdim_tÚpos_xÚpos_ys           r5   Úencode_1d_positionsz-Sam3SinePositionEmbedding.encode_1d_positionsi  si  € ð �d”j‘.ˆØ�d”j‘.ˆå”˜TÔ7½u¼{ÐSTÔS[Ð\Ñ\Ô\×_Ò_Ð`aÔ`gÑhÔhˆØÔ  Q¨%°1©*Ñ%5¸Ô8RÑ%RÑSˆà˜˜˜˜4˜Ô  5Ñ(ˆØ˜˜˜˜4˜Ô  5Ñ(ˆÝ”˜U 1 1 1 a d¨ d 7œ^×/Ò/Ñ1Ô1°5¸¸¸¸A¸D¸q¸D¸´>×3EÒ3EÑ3GÔ3GÐHÈaÐPÑPÔP×XÒXÐYZÑ[Ô[ˆÝ”˜U 1 1 1 a d¨ d 7œ^×/Ò/Ñ1Ô1°5¸¸¸¸A¸D¸q¸D¸´>×3EÒ3EÑ3GÔ3GÐHÈaÐPÑPÔP×XÒXÐYZÑ[Ô[ˆØ�eˆ|Ðr4   Úboxesc           	      ó*  — |                      d¦  «        dk    sJ d|j        › �¦   «         ‚t          j        | j        t          j        |j        ¬¦  «                             |j        ¦  «        }| j	        dt          j
        |dd¬¦  «        z  | j        z  z  }|dd…dd…d	f         | j        z  }|dd…dd…d
f         | j        z  }|dd…dd…df         | j        z  }|dd…dd…df         | j        z  }|dd…dd…df         |z  }|dd…dd…df         |z  }|dd…dd…df         |z  }	|dd…dd…df         |z  }
t          j        |dd…dd…d	dd…f                              ¦   «         |dd…dd…d
dd…f                              ¦   «         fd¬¦  «                             d¦  «        }t          j        |dd…dd…d	dd…f                              ¦   «         |dd…dd…d
dd…f                              ¦   «         fd¬¦  «                             d¦  «        }t          j        |	dd…dd…d	dd…f                              ¦   «         |	dd…dd…d
dd…f                              ¦   «         fd¬¦  «                             d¦  «        }	t          j        |
dd…dd…d	dd…f                              ¦   «         |
dd…dd…d
dd…f                              ¦   «         fd¬¦  «                             d¦  «        }
t          j        |||	|
fd¬¦  «        }|S )a:  
        Encode 4D box coordinates (x, y, w, h) for decoder conditioning using sine/cosine embeddings.

        Args:
            boxes: Box coordinates [batch_size, num_queries, 4] in (x, y, w, h) format

        Returns:
            Position embeddings [batch_size, num_queries, num_position_features*4]
        rh   rÝ   z4Expected 4D box coordinates (x, y, w, h), got shape r™  r   rß   rà   Nr   r   r   ri   )ro   rn   r0   rq   r‘  rš  rk   r"  rm   r’  ré   rÛ   r‹   rï   rî   rþ   rë   )r¤   r¡  r�  r›  rœ  Úw_embedÚh_embedrž  rŸ  Úpos_wÚpos_hÚposs               r5   Úencode_boxesz&Sam3SinePositionEmbedding.encode_boxes€  sk  € ð �zŠz˜"‰~Œ~ Ò"Ð"Ð"Ð$hÐ[`Ô[fÐ$hÐ$hÑ"Ô"Ð"Ý”˜TÔ7½u¼{ÐSXÔS_Ð`Ñ`Ô`×cÒcÐdiÔdoÑpÔpˆØÔ  Q­¬°5¸!È7Ð)SÑ)SÔ)SÑ%SÐVZÔVpÑ%pÑqˆà˜˜˜˜1˜1˜1˜a˜”. 4¤:Ñ-ˆØ˜˜˜˜1˜1˜1˜a˜”. 4¤:Ñ-ˆØ˜˜˜˜1˜1˜1˜a˜”. 4¤:Ñ-ˆØ˜˜˜˜1˜1˜1˜a˜”. 4¤:Ñ-ˆà˜˜˜˜1˜1˜1˜d˜
Ô# eÑ+ˆØ˜˜˜˜1˜1˜1˜d˜
Ô# eÑ+ˆØ˜˜˜˜1˜1˜1˜d˜
Ô# eÑ+ˆØ˜˜˜˜1˜1˜1˜d˜
Ô# eÑ+ˆå”˜U 1 1 1 a a a¨¨¨A¨ :Ô.×2Ò2Ñ4Ô4°e¸A¸A¸A¸q¸q¸qÀ!À$ÀQÀ$¸JÔ6G×6KÒ6KÑ6MÔ6MÐNÐTUÐVÑVÔV×^Ò^Ð_`ÑaÔaˆÝ”˜U 1 1 1 a a a¨¨¨A¨ :Ô.×2Ò2Ñ4Ô4°e¸A¸A¸A¸q¸q¸qÀ!À$ÀQÀ$¸JÔ6G×6KÒ6KÑ6MÔ6MÐNÐTUÐVÑVÔV×^Ò^Ð_`ÑaÔaˆÝ”˜U 1 1 1 a a a¨¨¨A¨ :Ô.×2Ò2Ñ4Ô4°e¸A¸A¸A¸q¸q¸qÀ!À$ÀQÀ$¸JÔ6G×6KÒ6KÑ6MÔ6MÐNÐTUÐVÑVÔV×^Ò^Ð_`ÑaÔaˆÝ”˜U 1 1 1 a a a¨¨¨A¨ :Ô.×2Ò2Ñ4Ô4°e¸A¸A¸A¸q¸q¸qÀ!À$ÀQÀ$¸JÔ6G×6KÒ6KÑ6MÔ6MÐNÐTUÐVÑVÔV×^Ò^Ð_`ÑaÔaˆåŒi˜  u¨eÐ4¸!Ð<Ñ<Ô<ˆàˆ
r4   rÝ   ©Úmaxsizern   rk   rm   Úmaskc           
      ó   — | \  }}	}
}|€wt          j        d|
dz   ||¬¦  «        d d d …d f                              ||
|¦  «        }t          j        d|dz   ||¬¦  «        d d d d …f                              ||
|¦  «        }n?|                     |¦  «        }|                     d¦  «        }|                     d¦  «        }|r6d}||d d …dd …d d …f         |z   z  |z  }||d d …d d …dd …f         |z   z  |z  }t          j        |t           j        |¬¦  «                             |¦  «        }|dt          j        |dd¬¦  «        z  |z  z  }|d d …d d …d d …d f         |z  }|d d …d d …d d …d f         |z  }t          j        |d d …d d …d d …dd d…f                              ¦   «         |d d …d d …d d …dd d…f          	                    ¦   «         fd	¬
¦  «         
                    d¦  «        }t          j        |d d …d d …d d …dd d…f                              ¦   «         |d d …d d …d d …dd d…f          	                    ¦   «         fd	¬
¦  «         
                    d¦  «        }t          j        ||fd¬
¦  «                             dddd¦  «        }|S )Nr   r™  r   g�íµ ÷Æ°>rh   rß   rà   r   rÝ   ri   r   )r0   rq   ru   r"  Úcumsumrš  ré   r‹   rï   rî   rþ   rë   r.  )rn   rk   rm   r‘  r“  rÛ   r’  r«  rz   r  r  r  rœ  r›  Ú
embed_maskrZ   r�  rž  rŸ  r§  s                       r5   Úbuild_sine_position_embeddingz7Sam3SinePositionEmbedding.build_sine_position_embedding¡  s  € ð (-Ñ$ˆ
�A�v˜uØˆ<õ ”l 1 f¨q¡j¸ÀfÐMÑMÔMÈdÐTUÐTUÐTUÐW[ÈmÔ\×cÒcØ˜F Eñô ˆGõ ”l 1 e¨a¡i°uÀVÐLÑLÔLÈTÐSWÐYZÐYZÐYZÈ]Ô[×bÒbØ˜F Eñô ˆGˆGð Ÿš ™œˆJØ ×'Ò'¨Ñ*Ô*ˆGØ ×'Ò'¨Ñ*Ô*ˆGØð 	CØˆCØ ¨¨¨¨B¨C¨C°°°¨Ô!3°cÑ!9Ñ:¸UÑBˆGØ ¨¨¨¨A¨A¨A¨r¨s¨s¨Ô!3°cÑ!9Ñ:¸UÑBˆGå”Ð2½%¼+ÈfÐUÑUÔU×XÒXÐY^Ñ_Ô_ˆØ ¥E¤I¨e°QÀgÐ$NÑ$NÔ$NÑ NÐQfÑ fÑgˆà˜˜˜˜1˜1˜1˜a˜a˜a ˜Ô&¨Ñ.ˆØ˜˜˜˜1˜1˜1˜a˜a˜a ˜Ô&¨Ñ.ˆÝ”˜U 1 1 1 a a a¨¨¨¨A¨D¨q¨D =Ô1×5Ò5Ñ7Ô7¸¸q¸q¸qÀ!À!À!ÀQÀQÀQÈÈÈ1È¸}Ô9M×9QÒ9QÑ9SÔ9SÐTÐZ[Ð\Ñ\Ô\×dÒdÐefÑgÔgˆÝ”˜U 1 1 1 a a a¨¨¨¨A¨D¨q¨D =Ô1×5Ò5Ñ7Ô7¸¸q¸q¸qÀ!À!À!ÀQÀQÀQÈÈÈ1È¸}Ô9M×9QÒ9QÑ9SÔ9SÐTÐZ[Ð\Ñ\Ô\×dÒdÐefÑgÔgˆÝŒi˜ ˜¨AÐ.Ñ.Ô.×6Ò6°q¸!¸QÀÑBÔBˆØˆ
r4   c           
      ób   — |                       |||| j        | j        | j        | j        |¦  «        S r•   )r¯  r‘  r“  rÛ   r’  )r¤   rn   rk   rm   r«  s        r5   r§   z!Sam3SinePositionEmbedding.forwardÍ  s9   € ð ×1Ò1Ø�6˜5 $Ô"<¸d¼nÈdÌjÐZ^ÔZjÐlpñ
ô 
ð 	
r4   )r�  r�  FN)FNr�  Nr•   )r+   r,   r-   r.   rù   rÏ   rç   r—   r0   r   r/   r   r¨  Ústaticmethodr   ÚSizerk   Ústrrm   r¯  r§   r©   rª   s   @r5   rŽ  rŽ  T  s  ø€ € € € € ðð ð &(Ø ØØ"ð=ð =à"ð=ð ð=ð ð	=ð
 �t‰|ð=ð =ð =ð =ð =ð =ð U¤\ð °e´lð ÀuÈUÌ\Ð[`Ô[gÐMgÔGhð ð ð ð ð. %¤,ð °5´<ð ð ð ð ðB Ø(Ð(°Ð3Ñ3Ô3ð  Ø"Ø Ø$(ð(ð (ØŒzð(à”˜sÑ"ð(ð Œ{ð(ð  #ð	(ð
 ð(ð �t‰|ð(ð ð(ð Œl˜TÑ!ð(ð 
Œð(ð (ð (ñ 4Ô3ñ „\ð(ð^ %)ð	
ð 	
àŒzð	
ð ”˜sÑ"ð	
ð Œ{ð		
ð
 Œl˜TÑ!ð	
ð 
Œð	
ð 	
ð 	
ð 	
ð 	
ð 	
ð 	
ð 	
r4   rŽ  c                   óP   ‡ — e Zd Zdededefˆ fd„Zdej        dej        fd„Zˆ xZ	S )ÚSam3FPNLayerÚin_channelsÚfpn_dimÚscale_factorc                 ó\  •— t          ¦   «                              ¦   «          || _        t          j        ¦   «         | _        |dk    rš| j                             t          j        ||dz  dd¬¦  «        ¦  «         | j                             t          j        ¦   «         ¦  «         | j                             t          j        |dz  |dz  dd¬¦  «        ¦  «         |dz  }n’|dk    r9| j                             t          j        ||dz  dd¬¦  «        ¦  «         |dz  }nS|dk    r|}nJ|dk    r1| j                             t          j	        dd¬¦  «        ¦  «         |}nt          d|› d	�¦  «        ‚t          j        ||d
¬¦  «        | _        t          j        ||dd
¬¦  «        | _        d S )Ng      @r   )r  r  rÝ   g       @rØ   r‰   zscale_factor=z is not supported yet.r   )r¶  Úout_channelsr  r   )r¶  rº  r  Úpadding)r–   r—   r¸  r›   r‚  Úscale_layersÚappendÚConvTranspose2dÚGELUÚ	MaxPool2dÚNotImplementedErrorr  Úproj1Úproj2)r¤   r¶  r·  r¸  Úintermediate_channelsr¥   s        €r5   r—   zSam3FPNLayer.__init__Ú  sÀ  ø€ Ý‰Œ×ÒÑÔÐØ(ˆÔõ œM™OœOˆÔà˜3ÒÐØÔ×$Ò$¥RÔ%7¸À[ÐTUÑEUÐcdÐmnÐ%oÑ%oÔ%oÑpÔpÐpØÔ×$Ò$¥R¤W¡Y¤YÑ/Ô/Ð/ØÔ×$Ò$¥RÔ%7¸ÀqÑ8HÈ+ÐYZÑJZÐhiÐrsÐ%tÑ%tÔ%tÑuÔuÐuØ$/°1Ñ$4Ð!Ð!Ø˜SÒ Ð ØÔ×$Ò$¥RÔ%7¸À[ÐTUÑEUÐcdÐmnÐ%oÑ%oÔ%oÑpÔpÐpØ$/°1Ñ$4Ð!Ð!Ø˜SÒ Ð Ø$/Ð!Ð!Ø˜SÒ Ð ØÔ×$Ò$¥R¤\¸aÈÐ%JÑ%JÔ%JÑKÔKÐKØ$/Ð!Ð!å%Ð&Z°lÐ&ZÐ&ZÐ&ZÑ[Ô[Ð[å”YÐ+@ÈwÐdeÐfÑfÔfˆŒ
Ý”Y¨7ÀÐVWÐabÐcÑcÔcˆŒ
ˆ
ˆ
r4   r@   r[   c                 óÌ   — |                      | j        j        j        ¦  «        }| j        D ]} ||¦  «        }Œ|                      |¦  «        }|                      |¦  «        }|S r•   )r"  rÂ  r#  rm   r¼  rÃ  )r¤   r@   r‹  s      r5   r§   zSam3FPNLayer.forwardô  sh   € Ø%×(Ò(¨¬Ô):Ô)@ÑAÔAˆØÔ&ð 	1ð 	1ˆEØ!˜E -Ñ0Ô0ˆMˆMàŸ
š
 =Ñ1Ô1ˆØŸ
š
 =Ñ1Ô1ˆàÐr4   )
r+   r,   r-   rù   rç   r—   r0   r   r§   r©   rª   s   @r5   rµ  rµ  Ù  s�   ø€ € € € € ðd Cð d°#ð dÀUð dð dð dð dð dð dð4 U¤\ð °e´lð ð ð ð ð ð ð ð r4   rµ  c                   ó‚   ‡ — e Zd Zdefˆ fd„Zdej        deeej        df         eej        df         f         fd„Zˆ xZ	S )ÚSam3VisionNeckr˜   c                 óè   •‡— t          ¦   «                              ¦   «          ‰| _        t          ‰j        dz  d¬¦  «        | _        t          j        ˆfd„‰j        D ¦   «         ¦  «        | _	        d S )Nr   T©r‘  r“  c                 óR   •— g | ]#}t          ‰j        j        ‰j        |¬ ¦  «        ‘Œ$S ))r¶  r·  r¸  )rµ  Úbackbone_configr|   Úfpn_hidden_size)r~  rÛ   r˜   s     €r5   r€  z+Sam3VisionNeck.__init__.<locals>.<listcomp>
  sK   ø€ ð ð ð ð õ Ø &Ô 6Ô BÈFÔLbÐqvðñ ô ðð ð r4   )
r–   r—   r˜   rŽ  rÌ  Úposition_encodingr›   r‚  Úscale_factorsÚ
fpn_layersr£   s    `€r5   r—   zSam3VisionNeck.__init__   s‰   øø€ Ý‰Œ×ÒÑÔÐØˆŒå!:Ø"(Ô"8¸AÑ"=Èð"
ñ "
ô "
ˆÔõ
 œ-ðð ð ð ð $Ô1ð	ñ ô ñ
ô 
ˆŒˆˆr4   r@   r[   .c                 ó    — d}d}| j         D ]?} ||¦  «        }||fz  }|                      |j        |j        |j        ¦  «        }||fz  }Œ@||fS )Nr3   )rÏ  rÍ  rn   rk   rm   )r¤   r@   r)   r*   Ú	fpn_layerÚ
fpn_outputÚpos_encs          r5   r§   zSam3VisionNeck.forward  sx   € ØÐØ "Ðàœð 	0ð 	0ˆIØ"˜ =Ñ1Ô1ˆJØ * Ñ.Ðà×,Ò,¨ZÔ-=¸zÔ?PÐR\ÔRbÑcÔcˆGØ! g ZÑ/Ð!Ð!à Ð"7Ð7Ð7r4   )
r+   r,   r-   r%   r—   r0   r   r/   r§   r©   rª   s   @r5   rÇ  rÇ  ÿ  s�   ø€ € € € € ð
Ð/ð 
ð 
ð 
ð 
ð 
ð 
ð$8 U¤\ð 8°e¸EÀ%Ä,ÐPSÐBSÔ<TÐV[Ð\aÔ\hÐjmÐ\mÔVnÐ<nÔ6oð 8ð 8ð 8ð 8ð 8ð 8ð 8ð 8r4   rÇ  zJ
    The vision model from Sam without any head or projection on top.
    )Úcustom_introc            	       ó|   ‡ — e Zd ZeZdZdefˆ fd„Zd„ Ze	 d	de	j
        dz  dee         deez  fd„¦   «         Zˆ xZS )
ÚSam3VisionModelr   r˜   c                 óä   •— t          ¦   «                              |¦  «         || _        t          j        |j        ¦  «        | _        t          |¦  «        | _        |  	                    ¦   «          d S r•   )
r–   r—   r˜   r   Úfrom_configrË  ÚbackbonerÇ  Úneckr†  r£   s     €r5   r—   zSam3VisionModel.__init__)  s\   ø€ Ý‰Œ×Ò˜Ñ Ô Ð ØˆŒÝ!Ô-¨fÔ.DÑEÔEˆŒÝ" 6Ñ*Ô*ˆŒ	à�ŠÑÔÐÐÐr4   c                 ó4   — | j                              ¦   «         S r•   )rÙ  rˆ  rø   s    r5   rˆ  z$Sam3VisionModel.get_input_embeddings1  s   € ØŒ}×1Ò1Ñ3Ô3Ð3r4   Nr±   r[   c                 ó¬  — |€t          d¦  «        ‚ | j        |fi |¤Ž}|j        }|j        d         }|j        d         | j        j        j        z  }|j        d         | j        j        j        z  }|                     |||d¦  «                             dddd¦  «        }|  	                    |¦  «        \  }	}
t          ||	|
|j        |j        ¬¦  «        S )	Nz You have to specify pixel_valuesr   rü   rh   r   r   r   )r8   r)   r*   r@   rA   )rå   rÙ  r8   rn   r˜   rË  r  rÌ   r.  rÚ  r(   r@   rA   )r¤   r   r±   Úbackbone_outputr@   rz   r  r  Úhidden_states_spatialr)   r*   s              r5   r§   zSam3VisionModel.forward4  sô   € ð ÐÝÐ?Ñ@Ô@Ð@à'˜$œ-¨Ð?Ð?¸Ð?Ð?ˆØ'Ô9ˆð #Ô(¨Ô+ˆ
ØÔ# BÔ'¨4¬;Ô+FÔ+QÑQˆØÔ" 2Ô&¨$¬+Ô*EÔ*PÑPˆØ -× 2Ò 2°:¸vÀuÈbÑ QÔ Q× YÒ YÐZ[Ð]^Ð`aÐcdÑ eÔ eÐØ37·9²9Ð=RÑ3SÔ3SÑ0ÐÐ0å&Ø+Ø/Ø"7Ø)Ô7Ø&Ô1ð
ñ 
ô 
ð 	
r4   r•   )r+   r,   r-   r%   rp  rr  r—   rˆ  r   r0   r1   r   r   r/   r(   r§   r©   rª   s   @r5   rÖ  rÖ     s»   ø€ € € € € ð $€LØ$€OðÐ/ð ð ð ð ð ð ð4ð 4ð 4ð ð 26ð
ð 
àÔ'¨$Ñ.ð
ð Ð+Ô,ð
ð 
Ð(Ñ	(ð	
ð 
ð 
ñ Ôð
ð 
ð 
ð 
ð 
r4   rÖ  c                   óL   ‡ — e Zd Zdefˆ fd„Zdededededee         f
d„Zˆ xZ	S )	ÚSam3GeometryEncoderLayerr˜   c                 ó°  •— t          ¦   «                              ¦   «          t          j        |j        ¦  «        | _        t          |¦  «        | _        t          j        |j	        ¦  «        | _	        t          |¦  «        | _
        t          j        |j        ¦  «        | _        t          |¦  «        | _        t          j        |j        ¦  «        | _        d S r•   )r–   r—   r›   rU  r|   rW  r¿   Ú	self_attnr    r¢   Ú
cross_attnrZ  r“   r[  Úlayer_norm3r£   s     €r5   r—   z!Sam3GeometryEncoderLayer.__init__Q  sœ   ø€ Ý‰Œ×ÒÑÔÐÝœ<¨Ô(:Ñ;Ô;ˆÔÝ& vÑ.Ô.ˆŒÝ”z &¤.Ñ1Ô1ˆŒå'¨Ñ/Ô/ˆŒÝœ<¨Ô(:Ñ;Ô;ˆÔå˜6‘?”?ˆŒÝœ<¨Ô(:Ñ;Ô;ˆÔÐÐr4   Úprompt_featsÚvision_featsÚvision_pos_encodingÚprompt_maskr±   c                 ó¦  — |}|                       |¦  «        } | j        d||||dœ|¤Ž\  }}|                      |¦  «        |z   }|}|                      |¦  «        }||z   }	 | j        d||	|dœ|¤Ž\  }}|                      |¦  «        |z   }|}|                      |¦  «        }|                      |¦  «        }|                      |¦  «        |z   }|S )N©r­   r®   r¯   r9   ©r­   r®   r¯   r3   ©rW  râ  r¢   rZ  rã  rä  r[  )
r¤   rå  ræ  rç  rè  r±   r`  r@   r  r®   s
             r5   r§   z Sam3GeometryEncoderLayer.forward]  s  € ð  ˆØ×(Ò(¨Ñ6Ô6ˆØ)˜4œ>ð 
Ø ]¸-ÐXcð
ð 
Øgmð
ð 
Ñˆ�qð Ÿš ]Ñ3Ô3°hÑ>ˆØ ˆØ×(Ò(¨Ñ7Ô7ˆØÐ0Ñ0ˆØ*˜4œ?Ðf°ÀCÈ|ÐfÐfÐ_eÐfÐfÑˆ�qØŸš ]Ñ3Ô3°hÑ>ˆØ ˆØ×(Ò(¨Ñ7Ô7ˆØŸš Ñ/Ô/ˆØŸš ]Ñ3Ô3°hÑ>ˆàÐr4   )
r+   r,   r-   r#   r—   r   r   r   r§   r©   rª   s   @r5   rà  rà  P  s�   ø€ € € € € ð
<Ð8ð 
<ð 
<ð 
<ð 
<ð 
<ð 
<ðàðð ðð $ð	ð
 ðð Ð+Ô,ðð ð ð ð ð ð ð r4   rà  c                   óô   ‡ — e Zd ZdZdefˆ fd„Zdej        dej        dej        dej        dej        f
d	„Zd
„ Z		 ddej        dej        dej        de
ej        df         de
ej        df         dz  f
d„Zˆ xZS )ÚSam3GeometryEncodera®  
    Encoder for geometric prompts (boxes).

    Boxes are encoded using three approaches:
     - Direct projection: linear projection from coordinate space to hidden_size
     - Pooling: pool features from the backbone at the specified location (ROI align for boxes)
     - Position encoding: use position encoding of the box center

    These encodings are combined additively and further processed with transformer layers.
    r˜   c                 óŠ  •‡— t          ¦   «                              ¦   «          ‰| _        ‰j        | _        ‰j        | _        t          ‰j        dz  d¬¦  «        | _        t          j        d| j        ¦  «        | _	        t          j        d| j        ¦  «        | _
        t          j        d| j        ¦  «        | _        t          j        | j        | j        | j        ¦  «        | _        t          j        | j        dz   | j        ¦  «        | _        t          j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        t          j        | j        ¦  «        | _        t          j        ˆfd„t+          ‰j        ¦  «        D ¦   «         ¦  «        | _        t          j        | j        ¦  «        | _        d S )Nr   TrÉ  r   rÝ   c                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r3   )rà  ©r~  r  r˜   s     €r5   r€  z0Sam3GeometryEncoder.__init__.<locals>.<listcomp>�  s"   ø€ Ð$hÐ$hÐ$hÈ!Õ%=¸fÑ%EÔ%EÐ$hÐ$hÐ$hr4   )r–   r—   r˜   r|   Úroi_sizerŽ  rÍ  r›   Ú	EmbeddingÚlabel_embedÚ	cls_embedrœ   Úboxes_direct_projectr  Úboxes_pool_projectÚboxes_pos_enc_projectrU  Úvision_layer_normÚ
final_projÚprompt_layer_normr‚  rƒ  Ú
num_layersr…  Úoutput_layer_normr£   s    `€r5   r—   zSam3GeometryEncoder.__init__„  sv  øø€ Ý‰Œ×ÒÑÔÐØˆŒØ!Ô-ˆÔØœˆŒå!:Ø"(Ô"4¸Ñ"9ÀTð"
ñ "
ô "
ˆÔõ œ<¨¨4Ô+;Ñ<Ô<ˆÔÝœ a¨Ô)9Ñ:Ô:ˆŒõ %'¤I¨a°Ô1AÑ$BÔ$BˆÔ!Ý"$¤)¨DÔ,<¸dÔ>NÐPTÔP]Ñ"^Ô"^ˆÔÝ%'¤Y¨tÔ/?À!Ñ/CÀTÔEUÑ%VÔ%VˆÔ"õ "$¤¨dÔ.>Ñ!?Ô!?ˆÔõ œ) DÔ$4°dÔ6FÑGÔGˆŒÝ!#¤¨dÔ.>Ñ!?Ô!?ˆÔõ ”mÐ$hÐ$hÐ$hÐ$hÍuÐU[ÔUfÑOgÔOgÐ$hÑ$hÔ$hÑiÔiˆŒÝ!#¤¨dÔ.>Ñ!?Ô!?ˆÔÐÐr4   Úcenter_xÚcenter_yr  r  r[   c                 óž   — | j                              ||¦  «        \  }}t          j        |||dd…df         |dd…df         fd¬¦  «        }|S )a�  
        Encode box coordinates by combining position-encoded centers with raw width/height.

        Args:
            center_x: 1D tensor of box center x coordinates
            center_y: 1D tensor of box center y coordinates
            width: 1D tensor of box widths
            height: 1D tensor of box heights

        Returns:
            Encoded box coordinates [N, embedding_dim]
        Nr   ri   )rÍ  r   r0   rë   )r¤   rþ  rÿ  r  r  rž  rŸ  r§  s           r5   Ú_encode_box_coordinatesz+Sam3GeometryEncoder._encode_box_coordinates   sZ   € ð Ô-×AÒAÀ(ÈHÑUÔU‰ˆˆuÝŒi˜  v¨a¨a¨a°¨g¤¸¸a¸a¸aÀ¸g¼ÐGÈQÐOÑOÔOˆØˆ
r4   c                 ó†  — |j         dd…         \  }}|j         dd…         \  }}|                      |¦  «        }	t          |¦  «        }
t          j        ||||g|
j        |
j        ¬¦  «        }|                     ddd¦  «        }|
|z  }
|j        t          j        k    rt          j	        n|j        }t          j                             |                     |¦  «        |
                     |¦  «                             d¦  «        | j        ¦  «                             |j        ¦  «        }|                      |¦  «        }|                     ||| j        ¦  «        }|	|z   }	|                     d¦  «        \  }}}}|                      |                     ¦   «         |                     ¦   «         |                     ¦   «         |                     ¦   «         ¦  «        }|                     |||j         d         ¦  «        }|                      |¦  «        }|	|z   }	|                      |                     ¦   «         ¦  «        }||	z   |fS )	z?Encode box prompts. Mask convention: True=valid, False=padding.Nr   rü   r™  r   rÝ   r   rh   )rn   rö  r‘   r0   Útensorrm   rk   rÌ   Úbfloat16Úfloat16ra  ÚopsÚ	roi_alignr"  rŠ   rò  r÷  r|   r  rþ   rø  rô  rè   )r¤   r¡  Ú
boxes_maskÚboxes_labelsÚvision_featuresrz   Ú	num_boxesr  r  Úboxes_embedÚ
boxes_xyxyrÛ   rm   Úsampled_featuresÚpooled_projectionrþ  rÿ  Ú	box_widthÚ
box_heightrÓ  Úpos_projectionrô  s                         r5   Ú_encode_boxesz!Sam3GeometryEncoder._encode_boxes³  s  € à %¤¨B¨Q¨B¤Ñˆ
�IØ'Ô-¨b¨c¨cÔ2‰ˆ�Ø×/Ò/°Ñ6Ô6ˆõ (¨Ñ.Ô.ˆ
Ý”˜e V¨U°FÐ;À:ÔCSÐ\fÔ\mÐnÑnÔnˆØ—
’
˜1˜a Ñ#Ô#ˆØ %Ñ'ˆ
ð "1Ô!6½%¼.Ò!HÐ!H•”�ÈoÔNcˆÝ&œ?×4Ò4Ø×Ò˜uÑ%Ô% z§}¢}°UÑ';Ô';×'BÒ'BÀ1Ñ'EÔ'EÀtÄ}ñ
ô 
ç
Š"ˆ_Ô"Ñ
#Ô
#ð 	ð !×3Ò3Ð4DÑEÔEÐØ-×2Ò2°:¸yÈ$ÔJZÑ[Ô[ÐØ!Ð$5Ñ5ˆð 5:·L²LÀÑ4DÔ4DÑ1ˆ�(˜I zØ×.Ò.Ø×ÒÑÔ × 0Ò 0Ñ 2Ô 2°I×4EÒ4EÑ4GÔ4GÈ×I[ÒI[ÑI]ÔI]ñ
ô 
ˆð —,’,˜z¨9°g´mÀBÔ6GÑHÔHˆØ×3Ò3°GÑ<Ô<ˆØ! NÑ2ˆð ×&Ò& |×'8Ò'8Ñ':Ô':Ñ;Ô;ˆØ˜[Ñ(¨*Ð4Ð4r4   NÚbox_embeddingsÚbox_maskÚ
box_labelsÚ	img_feats.Úimg_pos_embedsc                 óè  — |j         d         }|d         }|�|d         nt          j        |¦  «        }|                     d¦  «                             dd¦  «        }	|                     d¦  «                             dd¦  «        }
|d         }|                     dddd¦  «        }|                      |¦  «        }|                     dddd¦  «        }|                      ||||¦  «        \  }}| j        j	         
                    d| j        ¦  «                             d¦  «                             |dd¦  «        }t          j        |d|j        |j        ¬¦  «        }t#          ||||¦  «        \  }}|                      |                      |¦  «        ¦  «        }d}|�t)          | j        ||¬¦  «        }| j        D ]} |||	|
|¬	¦  «        }Œ|                      |¦  «        }t1          ||¬
¦  «        S )a8  
        Forward pass for encoding geometric prompts.

        Args:
            box_embeddings: Box coordinates in CxCyWH format [batch_size, num_boxes, 4]
            box_mask: Attention mask for boxes [batch_size, num_boxes]
            box_labels: Labels for boxes (positive/negative) [batch_size, num_boxes]
            img_feats: Image features from vision encoder
            img_pos_embeds: Optional position embeddings for image features

        Returns:
            Sam3GeometryEncoderOutput containing encoded geometry features and attention mask.
        r   rh   Nr   r   r   r™  )r˜   Úinputs_embedsr9   )rå  ræ  rç  rè  )r8   r9   )rn   r0   Ú
zeros_likerþ   r·   r.  rù  r  rõ  r#  rÌ   r|   Ú	unsqueezeru   rM  rm   rk   r‡   rû  rú  r   r˜   r…  rý  r7   )r¤   r  r  r  r  r  rz   ræ  Úvision_pos_embedsÚvision_feats_flatÚvision_pos_embeds_flatÚimg_feats_lastÚnormalized_img_featsÚprompt_embedsrè  rõ  Úcls_maskÚprompt_attention_maskr‹  s                      r5   r§   zSam3GeometryEncoder.forward×  s(  € ð* $Ô)¨!Ô,ˆ
ð ! ”}ˆØ2@Ð2L˜N¨2Ô.Ð.ÕRWÔRbÐcoÑRpÔRpÐØ(×0Ò0°Ñ3Ô3×=Ò=¸aÀÑCÔCÐØ!2×!:Ò!:¸1Ñ!=Ô!=×!GÒ!GÈÈ1Ñ!MÔ!MÐð # 2œˆØ'×/Ò/°°1°a¸Ñ;Ô;ˆØ#×5Ò5°nÑEÔEÐØ3×;Ò;¸A¸qÀ!ÀQÑGÔGÐà%)×%7Ò%7¸ÈÐR\Ð^rÑ%sÔ%sÑ"ˆ�{ð ”NÔ)×.Ò.¨q°$Ô2BÑCÔC×MÒMÈaÑPÔP×WÒWÐXbÐdfÐhjÑkÔkˆ	Ý”:˜j¨!°;Ô3DÈ[ÔM_Ð`Ñ`Ô`ˆÝ%<¸]ÈKÐYbÐdlÑ%mÔ%mÑ"ˆ�{à×.Ò.¨t¯ª¸}Ñ/MÔ/MÑNÔNˆð !%ÐØÐ"Ý$=Ø”{Ø+Ø*ð%ñ %ô %Ð!ð ”[ð 	ð 	ˆEØ!˜EØ*Ø.Ø$:Ø1ð	ñ ô ˆMˆMð ×.Ò.¨}Ñ=Ô=ˆå(Ø+Ø&ð
ñ 
ô 
ð 	
r4   r•   )r+   r,   r-   r.   r#   r—   r0   r   r  r  r/   r§   r©   rª   s   @r5   rî  rî  x  s,  ø€ € € € € ð	ð 	ð@Ð8ð @ð @ð @ð @ð @ð @ð8ØœðØ05´ðØEJÄ\ðØ[`Ô[gðà	Œðð ð ð ð&"5ð "5ð "5ðT ;?ðD
ð D
àœðD
ð ”,ðD
ð ”Lð	D
ð
 ˜œ sÐ*Ô+ðD
ð ˜eœl¨CÐ/Ô0°4Ñ7ðD
ð D
ð D
ð D
ð D
ð D
ð D
ð D
r4   rî  c                   óZ   ‡ — e Zd ZdZdefˆ fd„Z	 ddededededz  d	ee         f
d
„Z	ˆ xZ
S )ÚSam3DetrEncoderLayerz;DETR encoder layer with self-attention and cross-attention.r˜   c                 ó¾  •— t          ¦   «                              ¦   «          || _        t          j        |j        ¦  «        | _        t          |¦  «        | _        t          j	        |j
        ¦  «        | _
        t          |¦  «        | _        t          j        |j        ¦  «        | _        t          |¦  «        | _        t          j        |j        ¦  «        | _        d S r•   )r–   r—   r˜   r›   rU  r|   rW  r¿   râ  r    r¢   rã  rZ  r“   r[  rä  r£   s     €r5   r—   zSam3DetrEncoderLayer.__init__!  s£   ø€ Ý‰Œ×ÒÑÔÐØˆŒÝœ<¨Ô(:Ñ;Ô;ˆÔÝ& vÑ.Ô.ˆŒÝ”z &¤.Ñ1Ô1ˆŒå'¨Ñ/Ô/ˆŒÝœ<¨Ô(:Ñ;Ô;ˆÔå˜6‘?”?ˆŒÝœ<¨Ô(:Ñ;Ô;ˆÔÐÐr4   Nræ  rå  rç  Úprompt_cross_attn_maskr±   c                 ó¦  — |}|                       |¦  «        }||z   } | j        d|||dœ|¤Ž\  }}	|                      |¦  «        |z   }|}|                      |¦  «        } | j        d||||dœ|¤Ž\  }}	|                      |¦  «        |z   }|}|                      |¦  «        }|                      |¦  «        }|                      |¦  «        |z   }|S )a
  
        Forward pass for DETR encoder layer.

        Args:
            vision_feats: Vision features [batch_size, vision_len, hidden_size] (main hidden states)
            prompt_feats: Text prompt features [batch_size, text_len, hidden_size]
            vision_pos_encoding: Position encoding for vision [batch_size, vision_len, hidden_size]
            prompt_cross_attn_mask: Cross-attention mask for prompt features

        Returns:
            Updated vision features [batch_size, vision_len, hidden_size]
        rë  rê  r3   rì  )
r¤   ræ  rå  rç  r(  r±   r`  r@   Úhidden_states_with_posr  s
             r5   r§   zSam3DetrEncoderLayer.forward.  s&  € ð*  ˆØ×(Ò(¨Ñ6Ô6ˆØ!.Ð1DÑ!DÐØ)˜4œ>ð 
Ø(Ø&Øð
ð 
ð ð	
ð 
Ñˆ�qð Ÿš ]Ñ3Ô3°hÑ>ˆð !ˆØ×(Ò(¨Ñ7Ô7ˆà*˜4œ?ð 
ØØØØ1ð	
ð 
ð
 ð
ð 
Ñˆ�qð Ÿš ]Ñ3Ô3°hÑ>ˆð !ˆØ×(Ò(¨Ñ7Ô7ˆØŸš Ñ/Ô/ˆØŸš ]Ñ3Ô3°hÑ>ˆàÐr4   r•   )r+   r,   r-   r.   r"   r—   r   r   r   r§   r©   rª   s   @r5   r&  r&    s£   ø€ € € € € ØEÐEð<Ð4ð <ð <ð <ð <ð <ð <ð$ 15ð3ð 3àð3ð ð3ð $ð	3ð
 !'¨¡ð3ð Ð+Ô,ð3ð 3ð 3ð 3ð 3ð 3ð 3ð 3r4   r&  c                   ó:  ‡ — e Zd ZdZeedœZdefˆ fd„Zde	e
j                 de	e
j                 fd„Zee	 	 	 dde	e
j                 d	e
j        de	e
j                 dz  d
e
j        dz  de	eeef                  dz  dee         deez  fd„¦   «         ¦   «         Zˆ xZS )ÚSam3DetrEncodera  
    DETR-style encoder that processes multi-level vision features with text fusion.

    This encoder processes vision features from multiple levels (e.g., FPN features at different
    resolutions) and fuses them with text prompts through a stack of transformer encoder layers.
    rz  r˜   c                 ó  •‡— t          ¦   «                              ‰¦  «         ‰| _        ‰j        | _        t	          j        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        |  	                    ¦   «          d S )Nc                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r3   )r&  rñ  s     €r5   r€  z,Sam3DetrEncoder.__init__.<locals>.<listcomp>v  ó"   ø€ Ð$dÐ$dÐ$dÀaÕ%9¸&Ñ%AÔ%AÐ$dÐ$dÐ$dr4   )
r–   r—   r˜   r|   r›   r‚  rƒ  rü  r…  r†  r£   s    `€r5   r—   zSam3DetrEncoder.__init__q  sv   øø€ Ý‰Œ×Ò˜Ñ Ô Ð ØˆŒØ!Ô-ˆÔå”mÐ$dÐ$dÐ$dÐ$dÍ5ÐQWÔQbÑKcÔKcÐ$dÑ$dÔ$dÑeÔeˆŒà�ŠÑÔÐÐÐr4   r
  r  c                 ó2  — g }g }g }t          ||¦  «        D ]ª\  }}|j        dd…         \  }}	|                     ||	f¦  «         |                     d¦  «                             dd¦  «        }|                     d¦  «                             dd¦  «        }|                     |¦  «         |                     |¦  «         Œ«t          j        |d¬¦  «        }t          j        |d¬¦  «        }t          j        |t
          j        |j	        ¬¦  «        }|||fS )aÎ  
        Prepare multi-level vision features by flattening spatial dimensions and adding level embeddings.

        Args:
            vision_features: List of vision features at different levels [batch_size, channels, height, width]
            vision_pos_embeds: List of position embeddings for each level [batch_size, channels, height, width]

        Returns:
            Tuple containing flattened features, position embeddings, and spatial metadata
        rü   Nr   r   ri   r™  )
Úziprn   r½  rþ   r·   r0   rë   r  rè   rk   )
r¤   r
  r  Úfeatures_flattenedr=   r?   Úfeaturesr1  r  r  s
             r5   Ú_prepare_multilevel_featuresz,Sam3DetrEncoder._prepare_multilevel_featuresz  s2  € ð  ÐØ!ÐØˆå#& Ð8IÑ#JÔ#Jð 		3ð 		3ÑˆH�iØ$œN¨2¨3¨3Ô/‰MˆF�EØ×!Ò! 6¨5 /Ñ2Ô2Ð2ð  ×'Ò'¨Ñ*Ô*×4Ò4°Q¸Ñ:Ô:ˆHØ!×)Ò)¨!Ñ,Ô,×6Ò6°q¸!Ñ<Ô<ˆIà×%Ò% hÑ/Ô/Ð/Ø ×'Ò'¨	Ñ2Ô2Ð2Ð2õ #œYÐ'9¸qÐAÑAÔAÐÝ$œyÐ)=À1ÐEÑEÔEÐåœ n½E¼JÐOaÔOhÐiÑiÔiˆð Ø Øð
ð 	
r4   Nr>   Ú	text_maskÚspatial_sizesr±   r[   c                 ó^  — |d                               ¦   «         dk    r|d         j        d         n|d         j        d         }|�†t          |¦  «        D ]v\  }\  }	}
||                              |	|
|d¦  «                             dddd¦  «        ||<   ||                              |	|
|d¦  «                             dddd¦  «        ||<   Œw|                      ||¦  «        \  }}}d}|�t          | j        |||¬¦  «        }|}| j        D ]} ||f|||d	œ|¤Ž}Œt          ||||¬
¦  «        S )a)  
        Forward pass for the DETR encoder.

        Args:
            vision_features: List of vision features at different levels
            text_features: Text prompt features [batch_size, seq_len, hidden_size]
            vision_pos_embeds: Optional list of position embeddings for each level
            text_mask: Optional text padding mask [batch_size, seq_len]
            spatial_sizes: Optional list of (height, width) tuples for reshaping

        Returns:
            Sam3DETREncoderOutput containing encoded features and metadata.
        r   rÝ   r   Nrh   r   r   ©r˜   r  r9   rR   )rå  rç  r(  )r8   r=   r>   r?   )
rj   rn   Ú	enumeraterÒ   r.  r4  r   r˜   r…  r<   )r¤   r
  r>   r  r5  r6  r±   rz   r  r  r  r2  r=   r?   r(  r@   r‹  s                    r5   r§   zSam3DetrEncoder.forward¤  s´  € ð0 5DÀAÔ4F×4JÒ4JÑ4LÔ4LÐPQÒ4QÐ4Q�_ QÔ'Ô-¨aÔ0Ð0ÐWfÐghÔWiÔWoÐpqÔWrˆ
ð Ð$Ý&/°Ñ&>Ô&>ð wð wÑ"�‘?�F˜Eà%4°QÔ%7×%?Ò%?ÀÈÈzÐ[]Ñ%^Ô%^×%fÒ%fÐghÐjkÐmnÐpqÑ%rÔ%r� Ñ"Ø'8¸Ô';×'CÒ'CÀFÈEÐS]Ð_aÑ'bÔ'b×'jÒ'jÐklÐnoÐqrÐtuÑ'vÔ'vÐ! !Ñ$Ð$ð ×-Ò-¨oÐ?PÑQÔQñ		
ØØ Øð "&ÐØÐ Ý%>Ø”{Ø0Ø(Ø&3ð	&ñ &ô &Ð"ð +ˆØ”[ð 	ð 	ˆEØ!˜EØðà*Ø$8Ø'=ð	ð ð
 ðð ˆMˆMõ %Ø+Ø!5Ø'Ø)ð	
ñ 
ô 
ð 	
r4   )NNN)r+   r,   r-   r.   r&  r¿   rŒ  r"   r—   rT  r0   r   r4  r   r   r/   rù   r   r   r<   r§   r©   rª   s   @r5   r,  r,  d  sR  ø€ € € € € ðð ð .Ø#ðð Ðð
Ð4ð ð ð ð ð ð ð(
à˜eœlÔ+ð(
ð   ¤Ô-ð(
ð (
ð (
ð (
ðT  Øð
 8<Ø)-Ø6:ð=
ð =
à˜eœlÔ+ð=
ð ”|ð=
ð   ¤Ô-°Ñ4ð	=
ð
 ”< $Ñ&ð=
ð ˜E # s (œOÔ,¨tÑ3ð=
ð Ð+Ô,ð=
ð 
Ð&Ñ	&ð=
ð =
ð =
ñ „_ñ  Ôð=
ð =
ð =
ð =
ð =
r4   r,  c            	       óZ   ‡ — e Zd ZdZddedededefˆ fd„Zdej        d	ej        fd
„Zˆ xZ	S )ÚSam3DecoderMLPz/Simple 2 or 3-layer MLP for decoder components.r   Ú	input_dimÚ
hidden_dimÚ
output_dimrü  c                 óš  •— t          ¦   «                              ¦   «          |dk    r=t          j        ||¦  «        | _        t          j        ||¦  «        | _        d | _        d S |dk    rPt          j        ||¦  «        | _        t          j        ||¦  «        | _        t          j        ||¦  «        | _        d S t          d|› �¦  «        ‚)Nr   r   z"Only 2 or 3 layers supported, got )r–   r—   r›   rœ   Úlayer1Úlayer2Úlayer3rå   )r¤   r<  r=  r>  rü  r¥   s        €r5   r—   zSam3DecoderMLP.__init__é  s°   ø€ Ý‰Œ×ÒÑÔÐØ˜Š?ˆ?Ýœ) I¨zÑ:Ô:ˆDŒKÝœ) J°
Ñ;Ô;ˆDŒKØˆDŒKˆKˆKØ˜1Š_ˆ_Ýœ) I¨zÑ:Ô:ˆDŒKÝœ) J°
Ñ;Ô;ˆDŒKÝœ) J°
Ñ;Ô;ˆDŒKˆKˆKåÐNÀ*ÐNÐNÑOÔOÐOr4   rY   r[   c                 ó  — t          j        |                      |¦  «        ¦  «        }| j        �=t          j        |                      |¦  «        ¦  «        }|                      |¦  «        }n|                      |¦  «        }|S r•   )ÚFÚrelur@  rB  rA  )r¤   rY   s     r5   r§   zSam3DecoderMLP.forwardö  sa   € ÝŒF�4—;’;˜q‘>”>Ñ"Ô"ˆØŒ;Ð"Ý”�t—{’{ 1‘~”~Ñ&Ô&ˆAØ—’˜A‘”ˆAˆAà—’˜A‘”ˆAØˆr4   )r   )
r+   r,   r-   r.   rù   r—   r0   r   r§   r©   rª   s   @r5   r;  r;  æ  s–   ø€ € € € € Ø9Ð9ðPð P #ð P°3ð PÀCð PÐUXð Pð Pð Pð Pð Pð Pð˜œð ¨%¬,ð ð ð ð ð ð ð ð r4   r;  c                   óÂ   ‡ — e Zd ZdZdefˆ fd„Z	 	 ddej        dej        dej        dej        d	ej        d
ej        dz  dej        dz  dee	         dej        fd„Z
ˆ xZS )ÚSam3DetrDecoderLayerzYDETR decoder layer with self-attention, text cross-attention, and vision cross-attention.r˜   c                 óÖ  •— t          ¦   «                              ¦   «          || _        t          |¦  «        | _        t          j        |j        ¦  «        | _        t          j	        |j
        ¦  «        | _        t          |¦  «        | _        t          j        |j        ¦  «        | _        t          j	        |j
        ¦  «        | _        t          |¦  «        | _        t          j        |j        ¦  «        | _        t          j	        |j
        ¦  «        | _        t%          |¦  «        | _        t          j	        |j
        ¦  «        | _        t          j        |j        ¦  «        | _        d S r•   )r–   r—   r˜   r¿   râ  r›   r    r¢   Úself_attn_dropoutrU  r|   Úself_attn_layer_normÚtext_cross_attnÚtext_cross_attn_dropoutÚtext_cross_attn_layer_normÚvision_cross_attnÚvision_cross_attn_dropoutÚvision_cross_attn_layer_normr“   r[  Úmlp_layer_normÚmlp_dropoutr£   s     €r5   r—   zSam3DetrDecoderLayer.__init__  s  ø€ Ý‰Œ×ÒÑÔÐØˆŒÝ& vÑ.Ô.ˆŒÝ!#¤¨F¬NÑ!;Ô!;ˆÔÝ$&¤L°Ô1CÑ$DÔ$DˆÔ!å,¨VÑ4Ô4ˆÔÝ')¤z°&´.Ñ'AÔ'AˆÔ$Ý*,¬,°vÔ7IÑ*JÔ*JˆÔ'å!.¨vÑ!6Ô!6ˆÔÝ)+¬°F´NÑ)CÔ)CˆÔ&Ý,.¬L¸Ô9KÑ,LÔ,LˆÔ)å˜6‘?”?ˆŒÝ œl¨6Ô+=Ñ>Ô>ˆÔÝœ: f¤nÑ5Ô5ˆÔÐÐr4   Nr@   Ú	query_posr>   r
  rç  Útext_cross_attn_maskÚvision_cross_attn_maskr±   r[   c                 ó~  — t          j        |ddd¬¦  «        }|}	||z   }
 | j        d|
|
|ddœ|¤Ž\  }}|	|                      |¦  «        z   }|                      |¦  «        }|}	||z   }
 | j        d|
|||dœ|¤Ž\  }}|	|                      |¦  «        z   }|                      |¦  «        }|}	||z   }
||z   } | j        d|
|||dœ|¤Ž\  }}|	|  	                    |¦  «        z   }|  
                    |¦  «        }|}	|                      |¦  «        }|	|                      |¦  «        z   }|                      |¦  «        }|S )a  
        Forward pass for decoder layer.

        Args:
            hidden_states: Query features [batch_size, num_queries + 1, hidden_size] (includes presence token at position 0)
            query_pos: Query position embeddings [batch_size, num_queries, hidden_size]
            text_features: Text features [batch_size, seq_len, hidden_size]
            vision_features: Vision features [batch_size, height*width, hidden_size]
            vision_pos_encoding: Vision position encoding [batch_size, height*width, hidden_size]
            text_cross_attn_mask: Text cross-attention mask
            vision_cross_attn_mask: Vision cross-attention mask, already expanded for presence token

        Returns:
            Updated hidden states (including presence token at position 0)
        ©r   r   r   r   Úconstantr   ©Úmoder¯   Nrê  r3   )rD  r<  râ  rI  rJ  rK  rL  rM  rN  rO  rP  r[  rR  rQ  )r¤   r@   rS  r>   r
  rç  rT  rU  r±   r`  Úquery_with_posr¼   r  Úkey_with_poss                 r5   r§   zSam3DetrDecoderLayer.forward  sÑ  € õ6 ”E˜) \¸
È!ÐLÑLÔLˆ	ð !ˆØ&¨Ñ2ˆØ'˜œð 
Ø ØØØð	
ð 
ð
 ð
ð 
‰ˆ�Qð ! 4×#9Ò#9¸+Ñ#FÔ#FÑFˆØ×1Ò1°-Ñ@Ô@ˆð !ˆØ&¨Ñ2ˆà-˜Ô-ð 
Ø ØØØ/ð	
ð 
ð
 ð
ð 
‰ˆ�Qð ! 4×#?Ò#?ÀÑ#LÔ#LÑLˆØ×7Ò7¸ÑFÔFˆð !ˆØ&¨Ñ2ˆØ&Ð)<Ñ<ˆØ/˜Ô/ð 
Ø ØØ!Ø1ð	
ð 
ð
 ð
ð 
‰ˆ�Qð ! 4×#AÒ#AÀ+Ñ#NÔ#NÑNˆØ×9Ò9¸-ÑHÔHˆð !ˆØŸš Ñ/Ô/ˆØ  4×#3Ò#3°MÑ#BÔ#BÑBˆØ×+Ò+¨MÑ:Ô:ˆàÐr4   ©NN)r+   r,   r-   r.   r!   r—   r0   r   r   r   r§   r©   rª   s   @r5   rG  rG     sö   ø€ € € € € ØcÐcð6Ð4ð 6ð 6ð 6ð 6ð 6ð 6ð4 59Ø6:ðLð Là”|ðLð ”<ðLð ”|ð	Lð
 œðLð #œ\ðLð $œl¨TÑ1ðLð !&¤¨tÑ 3ðLð Ð+Ô,ðLð 
ŒðLð Lð Lð Lð Lð Lð Lð Lr4   rG  c                   ó¤  ‡ — e Zd ZdZeedœZdefˆ fd„Z e	d¬¦  «        de
j        de
j        d	e
j        d
e
j        dee
j        e
j        f         f
d„¦   «         Zde
j        dee
j        e
j        f         de
j        fd„Zee	 	 dde
j        de
j        de
j        de
j        dz  de
j        dz  dee         deez  fd„¦   «         ¦   «         Zˆ xZS )ÚSam3DetrDecodera  
    DETR-style decoder with box refinement and presence token.

    Simplified version that assumes:
    - Box refinement is always enabled
    - Intermediate outputs are always returned
    - BoxRPB (relative position bias) with log-scale encoding
    - Presence token is used
    rz  r˜   c                 óè  •‡— t          ¦   «                              ‰¦  «         ‰| _        ‰j        | _        t	          j        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        t	          j	        ‰j        ¦  «        | _
        t          ‰j        ‰j        dd¦  «        | _        t	          j        ‰j        ‰j        ¦  «        | _        t	          j        ‰j        d¦  «        | _        t	          j        d‰j        ¦  «        | _        t          ‰j        ‰j        dd¦  «        | _        t	          j	        ‰j        ¦  «        | _        d| _        t          d‰j        z  ‰j        ‰j        d¦  «        | _        t          d‰j        ‰j        d¦  «        | _        t          d‰j        ‰j        d¦  «        | _        t3          ‰j        dz  d¬¦  «        | _        |                      ¦   «          d S )	Nc                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r3   )rG  rñ  s     €r5   r€  z,Sam3DetrDecoder.__init__.<locals>.<listcomp>}  r/  r4   rÝ   r   r   g      $@r   FrÉ  )r–   r—   r˜   r|   r›   r‚  rƒ  rü  r…  rU  rý  r;  Úbox_headró  Únum_queriesÚquery_embedÚreference_pointsÚpresence_tokenÚpresence_headÚpresence_layer_normÚclamp_presence_logit_max_valÚref_point_headrÂ   Úbox_rpb_embed_xÚbox_rpb_embed_yrŽ  rÍ  r†  r£   s    `€r5   r—   zSam3DetrDecoder.__init__u  s²  øø€ õ 	‰Œ×Ò˜Ñ Ô Ð ØˆŒØ!Ô-ˆÔå”mÐ$dÐ$dÐ$dÐ$dÍ5ÐQWÔQbÑKcÔKcÐ$dÑ$dÔ$dÑeÔeˆŒå!#¤¨fÔ.@Ñ!AÔ!AˆÔå& vÔ'9¸6Ô;MÈqÐRSÑTÔTˆŒåœ<¨Ô(:¸FÔ<NÑOÔOˆÔÝ "¤¨VÔ-?ÀÑ CÔ CˆÔå œl¨1¨fÔ.@ÑAÔAˆÔÝ+¨FÔ,>ÀÔ@RÐTUÐWXÑYÔYˆÔÝ#%¤<°Ô0BÑ#CÔ#CˆÔ Ø,0ˆÔ)å,¨Q°Ô1CÑ-CÀVÔEWÐY_ÔYkÐmnÑoÔoˆÔå-¨a°Ô1CÀVÔE_ÐabÑcÔcˆÔÝ-¨a°Ô1CÀVÔE_ÐabÑcÔcˆÔå!:Ø"(Ô"4¸Ñ"9ÀUð"
ñ "
ô "
ˆÔð 	�ŠÑÔÐÐÐr4   r   r©  r  r  rm   rk   r[   c                 óv   — t          j        d|||¬¦  «        |z  }t          j        d|||¬¦  «        |z  }||fS )z%Generate normalized coordinate grids.r   rl   )r0   rq   )r¤   r  r  rm   rk   Úcoords_hÚcoords_ws          r5   Ú_get_coordszSam3DetrDecoder._get_coords–  sI   € õ
 ”<  6°&ÀÐFÑFÔFÈÑOˆÝ”<  5°¸uÐEÑEÔEÈÑMˆØ˜Ð!Ð!r4   rF   Úspatial_shapec                 óv  — |\  }}t          |¦  «        }|j        \  }}}|                      |||j        |j        ¬¦  «        \  }	}
|	                     ddd¦  «        |                     ddd¦  «        dd…dd…ddd…f         z
  }|                     ||dd¦  «        }|
                     ddd¦  «        |                     ddd¦  «        dd…dd…ddd…f         z
  }|                     ||dd¦  «        }|d	z  }t          j        |¦  «        t          j	        t          j
        |¦  «        d
z   ¦  «        z  t          j	        d	¦  «        z  }|d	z  }t          j        |¦  «        t          j	        t          j
        |¦  «        d
z   ¦  «        z  t          j	        d	¦  «        z  }|                      |¦  «        }|                      |¦  «        }|                     d¦  «        |                     d¦  «        z   }|                     dd¦  «        }|                     dddd¦  «                             ¦   «         }|S )aÓ  
        Compute box relative position bias (RPB) matrix using log-scale encoding.
        RPB helps the decoder attend to relevant spatial locations based on predicted box positions.

        Args:
            reference_boxes: Reference boxes [batch_size, num_queries, 4] in sigmoid space
            spatial_shape: (height, width) of the vision features as tensors

        Returns:
            RPB matrix [batch_size, num_heads, num_queries, height*width]
        r™  r   rh   rÝ   Nr   r   r   é   rØ   )r‘   rn   rp  rm   rk   rÌ   rÒ   r0   ÚsignÚlog2Úabsr•  rk  rl  r  rþ   r.  rº   )r¤   rF   rq  r  r  r  rz   rc  r  rn  ro  Údeltas_yÚdeltas_xÚdeltas_x_logÚdeltas_y_logÚ
rpb_matrixs                   r5   Ú_get_rpb_matrixzSam3DetrDecoder._get_rpb_matrixŸ  sH  € ð &‰ˆ�Ý'¨Ñ8Ô8ˆ
Ø%/Ô%5Ñ"ˆ
�K ð "×-Ò-Ø�E Ô!6¸Ô?Uð .ñ 
ô 
Ñˆ�(ð
 —=’=  B¨Ñ*Ô*¨Z×-?Ò-?ÀÀAÀqÑ-IÔ-IÈ!È!È!ÈQÈQÈQÐPQÐRSÐTUÐPUÈ+Ô-VÑVˆØ—=’= ¨[¸"¸aÑ@Ô@ˆØ—=’=  B¨Ñ*Ô*¨Z×-?Ò-?ÀÀAÀqÑ-IÔ-IÈ!È!È!ÈQÈQÈQÐPQÐRSÐTUÐPUÈ+Ô-VÑVˆØ—=’= ¨[¸"¸aÑ@Ô@ˆð   !‘|ˆÝ”z ,Ñ/Ô/µ%´*½U¼YÀ|Ñ=TÔ=TÐWZÑ=ZÑ2[Ô2[Ñ[Õ^bÔ^gÐhiÑ^jÔ^jÑjˆØ !‘|ˆÝ”z ,Ñ/Ô/µ%´*½U¼YÀ|Ñ=TÔ=TÐWZÑ=ZÑ2[Ô2[Ñ[Õ^bÔ^gÐhiÑ^jÔ^jÑjˆð ×'Ò'¨Ñ5Ô5ˆØ×'Ò'¨Ñ5Ô5ˆð ×'Ò'¨Ñ*Ô*¨X×-?Ò-?Øñ.
ô .
ñ 
ˆ
ð  ×'Ò'¨¨1Ñ-Ô-ˆ
Ø×'Ò'¨¨1¨a°Ñ3Ô3×>Ò>Ñ@Ô@ˆ
ØÐr4   Nr
  r>   rç  r5  r?   r±   c                 ó~  — |j         d         }| j        j                             d¦  «                             |dd¦  «        }| j        j                             d¦  «                             |dd¦  «        }	|	                     ¦   «         }	| j        j                             d¦  «                             |dd¦  «        }
t          j	        |
|gd¬¦  «        }d}|�t          | j        |||¬¦  «        }g }|	g}g }| j        D �]ç}|	                     d¦  «        }| j                             |dd…dd…ddd…f         ¦  «        }|                      |¦  «        }d}|�O|j         d         dk    r>|d         |d	         f}|                      |	|¦  «        }t#          j        |d
dd¬¦  «        } ||f||||||dœ|¤Ž}|dd…dd…f         }t'          |	¦  «        }|                      |                      |¦  «        ¦  «        }||z                        ¦   «         }|                     ¦   «         }	|                     |                      |¦  «        ¦  «         |                     |¦  «         |dd…dd…f         }|                      |                      |¦  «        ¦  «                             d¦  «        }|                     | j         | j        ¬¦  «        }|                     |¦  «         �Œét          j        |¦  «        }t          j        |dd…         ¦  «        }t          j        |¦  «        }t=          |||¬¦  «        S )a@  
        Forward pass for the DETR decoder.

        Args:
            vision_features: Vision features [batch_size, height*width, hidden_size]
            text_features: Text features [batch_size, seq_len, hidden_size]
            vision_pos_encoding: Vision position encoding [batch_size, height*width, hidden_size]
            text_mask: Text padding mask [batch_size, seq_len] where True=valid, False=padding
            spatial_shapes: Spatial shapes [num_levels, 2]

        Returns:
            Sam3DETRDecoderOutput containing decoder outputs from all layers.
        r   rh   r   ri   Nr8  r   )r   r   )r   r   rW  rX  rY  )rS  r>   r
  rç  rT  rU  r]   )rE   rF   rG   )rn   rd  r#  r  ru   re  Úsigmoidrf  r0   rë   r   r˜   r…  rÍ  r¨  rj  r|  rD  r<  re   rb  rý  Údetachr½  rg  rh  Úsqueezera   ri  r‹   rD   )r¤   r
  r>   rç  r5  r?   r±   rz   Úquery_embedsrF   rf  r@   rT  Úintermediate_outputsÚintermediate_boxesÚintermediate_presence_logitsr‹  Úreference_points_inputÚquery_sine_embedrS  rU  rq  r{  Úquery_hidden_statesÚreference_boxes_before_sigmoidÚdelta_boxesÚnew_reference_boxesÚpresence_hiddenrG   s                                r5   r§   zSam3DetrDecoder.forwardÎ  s®  € ð0 %Ô*¨1Ô-ˆ
àÔ'Ô.×8Ò8¸Ñ;Ô;×BÒBÀ:ÈrÐSUÑVÔVˆØÔ/Ô6×@Ò@ÀÑCÔC×JÒJÈ:ÐWYÐ[]Ñ^Ô^ˆØ)×1Ò1Ñ3Ô3ˆØÔ,Ô3×=Ò=¸aÑ@Ô@×GÒGÈ
ÐTVÐXZÑ[Ô[ˆõ œ	 >°<Ð"@ÀaÐHÑHÔHˆà#ÐØÐ Ý#<Ø”{Ø+Ø(Ø&3ð	$ñ $ô $Ð ð  "ÐØ-Ð.ÐØ')Ð$à”[ð +	Añ +	AˆEà%4×%>Ò%>¸qÑ%AÔ%AÐ"Ø#Ô5×BÒBÐCYÐZ[ÐZ[ÐZ[Ð]^Ð]^Ð]^Ð`aÐcdÐcdÐcdÐZdÔCeÑfÔfÐØ×+Ò+Ð,<Ñ=Ô=ˆIð &*Ð"ØÐ)¨nÔ.BÀ1Ô.EÈÒ.JÐ.JØ!/°Ô!5°~ÀdÔ7KÐ L�Ø!×1Ò1°/À=ÑQÔQ�
å)*¬¨z¸<ÈjÐ`aÐ)bÑ)bÔ)bÐ&à!˜EØð	à#Ø+Ø /Ø$7Ø%9Ø'=ð	ð 	ð ð	ð 	ˆMð #0°°°°1°2°2°Ô"6Ðõ .=¸_Ñ-MÔ-MÐ*ØŸ-š-¨×(>Ò(>Ð?RÑ(SÔ(SÑTÔTˆKØ#.Ð1OÑ#O×"XÒ"XÑ"ZÔ"ZÐØ1×8Ò8Ñ:Ô:ˆOà ×'Ò'¨×(>Ò(>Ð?RÑ(SÔ(SÑTÔTÐTØ×%Ò%Ð&9Ñ:Ô:Ð:ð ,¨A¨A¨A¨r°¨r¨EÔ2ˆOØ"×0Ò0°×1IÒ1IÈ/Ñ1ZÔ1ZÑ[Ô[×cÒcÐdfÑgÔgˆOØ-×3Ò3ØÔ6Ð6¸DÔ<]ð 4ñ ô ˆOð )×/Ò/°Ñ@Ô@Ð@Ñ@õ  %œ{Ð+?Ñ@Ô@ÐÝ"œ[Ð);¸C¸R¸CÔ)@ÑAÔAÐÝ',¤{Ð3OÑ'PÔ'PÐ$å$Ø';Ø.Ø8ð
ñ 
ô 
ð 	
r4   r]  )r+   r,   r-   r.   rG  r¿   rŒ  r!   r—   r   r0   r   rm   rk   r/   rp  r|  r   r   r   r   rD   r§   r©   rª   s   @r5   r_  r_  e  sÀ  ø€ € € € € ðð ð .Ø#ðð Ðð
à%ðð ð ð ð ð ðB )Ð(°Ð3Ñ3Ô3ð"Ø”lð"Ø+0¬<ð"Ø@EÄð"ØUZÔUað"à	ˆuŒ|˜Uœ\Ð)Ô	*ð"ð "ð "ñ 4Ô3ð"ð-Ø$œ|ð-Ø<AÀ%Ä,ÐPUÔP\ÐB\Ô<]ð-à	Œð-ð -ð -ð -ð^  Øð *.Ø.2ðc
ð c
àœðc
ð ”|ðc
ð #œ\ð	c
ð
 ”< $Ñ&ðc
ð œ tÑ+ðc
ð Ð+Ô,ðc
ð 
Ð&Ñ	&ðc
ð c
ð c
ñ „_ñ  Ôðc
ð c
ð c
ð c
ð c
r4   r_  c            	       óª   ‡ — e Zd ZdZdefˆ fd„Zdej        dej        dz  dej        fd„Z	 dd	ej        dej        dej        dz  dej        fd
„Z	ˆ xZ
S )ÚSam3DotProductScoringzÆ
    Computes classification scores by computing dot product between projected decoder queries and pooled text features.
    This is used to determine confidence/presence scores for each query.
    r˜   c                 ó  •— t          ¦   «                              ¦   «          || _        |j        j        }|j        j        }t          ||j        j        |d¬¦  «        | _        t          j	        |j        j
        ¦  «        | _        t          j        |¦  «        | _        t          j        ||¦  «        | _        t          j        ||¦  «        | _        t#          dt%          j        |¦  «        z  ¦  «        | _        d| _        d| _        d S )Nr   )r<  r=  r>  rü  rØ   Tg      (@)r–   r—   r˜   Údetr_decoder_configr|   r;  r�   Útext_mlpr›   r    r¢   Útext_mlp_dropoutrU  Útext_mlp_out_normrœ   Ú	text_projÚ
query_projrç   ÚnpÚsqrtrÛ   Úclamp_logitsÚclamp_max_val)r¤   r˜   r|   Úprojection_dimr¥   s       €r5   r—   zSam3DotProductScoring.__init__<  sê   ø€ Ý‰Œ×ÒÑÔÐØˆŒØÔ0Ô<ˆØÔ3Ô?ˆå&Ø!ØÔ1ÔCØ"Øð	
ñ 
ô 
ˆŒõ !#¤
¨6Ô+EÔ+MÑ NÔ NˆÔÝ!#¤¨kÑ!:Ô!:ˆÔõ œ ;°Ñ?Ô?ˆŒÝœ) K°Ñ@Ô@ˆŒõ ˜3¥¤¨Ñ!8Ô!8Ñ8Ñ9Ô9ˆŒ
ð !ˆÔØ!ˆÔÐÐr4   r>   r5  Nr[   c                 ó  — |€|                      d¬¦  «        S |                     |j        ¦  «                             d¦  «        }|                     d¬¦  «                             d¬¦  «        }||z                       d¬¦  «        |z  }|S )a<  
        Mean pool text features, accounting for padding.

        Args:
            text_features: [batch_size, seq_len, hidden_size]
            text_mask: [batch_size, seq_len] where True indicates valid tokens, False indicates padding

        Returns:
            pooled_text: [batch_size, hidden_size]
        Nr   ri   rh   rØ   r`   )ri  r"  rm   r  rp   ra   )r¤   r>   r5  Úis_validÚ	num_validÚpooled_texts         r5   Ú_pool_text_featuresz)Sam3DotProductScoring._pool_text_featuresV  s‘   € ð Ðà ×%Ò%¨!Ð%Ñ,Ô,Ð,à—<’< Ô 3Ñ4Ô4×>Ò>¸rÑBÔBˆð —L’L Q�LÑ'Ô'×-Ò-°#Ð-Ñ6Ô6ˆ	ð % xÑ/×4Ò4¸Ð4Ñ;Ô;¸iÑGˆàÐr4   rP   c                 óò  — |}|                       |¦  «        }|                      |¦  «        }||z   }|                      |¦  «        }|                      ||¦  «        }|                      |¦  «        }|                      |¦  «        }|                     d¦  «        }t          j        ||                     d¦  «        ¦  «        }|| j	        z  }| j
        r"|                     | j         | j        ¬¦  «        }|S )a  
        Compute classification scores via dot product.

        Args:
            decoder_hidden_states: [num_layers, batch_size, num_queries, hidden_size]
            text_features: [batch_size, seq_len, hidden_size]
            text_mask: [batch_size, seq_len] where True=valid, False=padding

        Returns:
            scores: [num_layers, batch_size, num_queries, 1]
        rh   r   r]   )r�  r‘  r’  rž  r“  r”  r  r0   r¶   rÛ   r—  ra   r˜  )	r¤   rP   r>   r5  Úorig_text_featuresr�  Ú	proj_textÚproj_queriesÚscoress	            r5   r§   zSam3DotProductScoring.forwardo  sñ   € ð" +ÐØŸš mÑ4Ô4ˆØ×-Ò-¨mÑ<Ô<ˆØ%Ð(:Ñ:ˆØ×.Ò.¨}Ñ=Ô=ˆà×.Ò.¨}¸iÑHÔHˆà—N’N ;Ñ/Ô/ˆ	Ø—’Ð'<Ñ=Ô=ˆà×'Ò'¨Ñ+Ô+ˆ	Ý”˜l¨I×,?Ò,?ÀÑ,BÔ,BÑCÔCˆØ˜$œ*Ñ$ˆØÔð 	SØ—\’\ tÔ'9Ð&9¸tÔ?Q�\ÑRÔRˆFàˆr4   r•   )r+   r,   r-   r.   r    r—   r0   r   rž  r§   r©   rª   s   @r5   r�  r�  6  sÖ   ø€ € € € € ðð ð
"˜zð "ð "ð "ð "ð "ð "ð4°´ð È%Ì,ÐY]ÑJ]ð ÐbgÔbnð ð ð ð ð: *.ð	"ð "à$œ|ð"ð ”|ð"ð ”< $Ñ&ð	"ð
 
Œð"ð "ð "ð "ð "ð "ð "ð "r4   r�  c                   óL   ‡ — e Zd ZdZdefˆ fd„Zdej        dej        fd„Zˆ xZ	S )ÚSam3MaskEmbedderzh
    MLP that embeds object queries for mask prediction.
    Similar to MaskFormer's mask embedder.
    r˜   c                 ó>  •— t          ¦   «                              ¦   «          || _        |j        }t	          j        t	          j        ||¦  «        t	          j        ||¦  «        t	          j        ||¦  «        g¦  «        | _        t	          j        ¦   «         | _	        d S r•   )
r–   r—   r˜   r|   r›   r‚  rœ   r…  ÚReLUÚ
activation©r¤   r˜   r|   r¥   s      €r5   r—   zSam3MaskEmbedder.__init__š  s€   ø€ Ý‰Œ×ÒÑÔÐØˆŒØÔ(ˆå”må”	˜+ {Ñ3Ô3Ý”	˜+ {Ñ3Ô3Ý”	˜+ {Ñ3Ô3ðñ
ô 
ˆŒõ œ'™)œ)ˆŒˆˆr4   Úqueriesr[   c                 ó´   — |}t          | j        ¦  «        D ]@\  }} ||¦  «        }|t          | j        ¦  «        dz
  k     r|                      |¦  «        }ŒA|S )z¹
        Args:
            queries: Query embeddings [batch_size, num_queries, hidden_size]

        Returns:
            Mask embeddings [batch_size, num_queries, hidden_size]
        r   )r9  r…  Úlenr¨  )r¤   rª  r@   r  r‹  s        r5   r§   zSam3MaskEmbedder.forward¨  sg   € ð  ˆÝ! $¤+Ñ.Ô.ð 	?ð 	?‰HˆAˆuØ!˜E -Ñ0Ô0ˆMØ•3�t”{Ñ#Ô# aÑ'Ò'Ð'Ø $§¢°Ñ >Ô >�øØÐr4   )
r+   r,   r-   r.   r$   r—   r0   r   r§   r©   rª   s   @r5   r¥  r¥  ”  su   ø€ € € € € ðð ð
$Ð4ð $ð $ð $ð $ð $ð $ð˜uœ|ð °´ð ð ð ð ð ð ð ð r4   r¥  c                   óX   ‡ — e Zd ZdZdefˆ fd„Zdeej                 dej        fd„Z	ˆ xZ
S )ÚSam3PixelDecoderz€
    Feature Pyramid Network (FPN) decoder that generates pixel-level features.
    Inspired by MaskFormer's pixel decoder.
    r˜   c                 óJ  •‡— t          ¦   «                              ¦   «          || _        |j        Š|j        }t          j        ˆfd„t          |¦  «        D ¦   «         ¦  «        | _        t          j        ˆfd„t          |¦  «        D ¦   «         ¦  «        | _	        ‰| _
        d S )Nc           	      óB   •— g | ]}t          j        ‰‰d dd¬¦  «        ‘ŒS )r   r   )r  r  r»  )r›   r  ©r~  r  r|   s     €r5   r€  z-Sam3PixelDecoder.__init__.<locals>.<listcomp>Æ  s?   ø€ ð ð ð àõ ”	˜+ {ÀÈ!ÐUVÐWÑWÔWðð ð r4   c                 ó:   •— g | ]}t          j        d ‰¦  «        ‘ŒS )rs  )r›   Ú	GroupNormr±  s     €r5   r€  z-Sam3PixelDecoder.__init__.<locals>.<listcomp>Ë  s%   ø€ Ð#gÐ#gÐ#gÀQ¥B¤L°°KÑ$@Ô$@Ð#gÐ#gÐ#gr4   )r–   r—   r˜   r|   Únum_upsampling_stagesr›   r‚  rƒ  Úconv_layersÚnormsrº  )r¤   r˜   r´  r|   r¥   s      @€r5   r—   zSam3PixelDecoder.__init__¾  s¶   øø€ Ý‰Œ×ÒÑÔÐØˆŒØÔ(ˆØ &Ô <Ðõ œ=ðð ð ð åÐ4Ñ5Ô5ðñ ô ñ
ô 
ˆÔõ ”]Ð#gÐ#gÐ#gÐ#gÍ%ÐPeÑJfÔJfÐ#gÑ#gÔ#gÑhÔhˆŒ
à'ˆÔÐÐr4   Úbackbone_featuresr[   c                 ó<  — |d         }t          t          |dd…         ¦  «        ¦  «        D ]n\  }}t          j        ||j        dd…         d¬¦  «        }||z   } | j        |         |¦  «        } | j        |         |¦  «        }t          j        |¦  «        }Œo|S )aA  
        Args:
            backbone_features: List of backbone features [batch_size, hidden_size, H_i, W_i]
                              from low to high resolution (assumes already projected to hidden_size)

        Returns:
            Pixel embeddings [batch_size, hidden_size, H, W] at the finest resolution
        rh   Nrü   Únearest)ro   rZ  )r9  ÚreversedrD  Úinterpolatern   rµ  r¶  rE  )r¤   r·  Úprev_fpnÚ	layer_idxÚbackbone_feats        r5   r§   zSam3PixelDecoder.forwardÏ  s°   € ð % RÔ(ˆå(1µ(Ð;LÈSÈbÈSÔ;QÑ2RÔ2RÑ(SÔ(Sð 
	(ð 
	(Ñ$ˆI�}å”} X°MÔ4GÈÈÈÔ4LÐS\Ð]Ñ]Ô]ˆHð   -Ñ/ˆHð 3�tÔ'¨	Ô2°8Ñ<Ô<ˆHØ,�t”z )Ô,¨XÑ6Ô6ˆHÝ”v˜hÑ'Ô'ˆHˆHàˆr4   )r+   r,   r-   r.   r$   r—   rT  r0   r   r§   r©   rª   s   @r5   r®  r®  ¸  sz   ø€ € € € € ðð ð
(Ð4ð (ð (ð (ð (ð (ð (ð"¨¨e¬lÔ);ð ÀÄð ð ð ð ð ð ð ð r4   r®  c                   ó  ‡ — e Zd ZdZdeiZdefˆ fd„Zee		 	 dde
j        dee
j                 de
j        d	e
j        dz  d
e
j        dz  dee         deez  fd„¦   «         ¦   «         Zdee
j                 de
j        de
j        fd„Zˆ xZS )ÚSam3MaskDecoderzÂ
    Mask decoder that combines object queries with pixel-level features to predict instance masks.
    Also produces a semantic segmentation output and supports cross-attention to prompts.
    rA   r˜   c                 ó  •— t          ¦   «                              |¦  «         || _        |j        }t	          |¦  «        | _        t          |¦  «        | _        t          j	        | j        j
        |d¬¦  «        | _        t          j	        | j        j
        dd¬¦  «        | _        t          |¦  «        | _        t          j        |¦  «        | _        t          j        |j        ¦  «        | _        |                      ¦   «          d S )Nr   )r  )r–   r—   r˜   r|   r®  Úpixel_decoderr¥  Úmask_embedderr›   r  rº  Úinstance_projectionÚsemantic_projectionr¿   Úprompt_cross_attnrU  Úprompt_cross_attn_normr    r¢   Úprompt_cross_attn_dropoutr†  r©  s      €r5   r—   zSam3MaskDecoder.__init__ô  sÞ   ø€ Ý‰Œ×Ò˜Ñ Ô Ð ØˆŒØÔ(ˆõ .¨fÑ5Ô5ˆÔõ .¨fÑ5Ô5ˆÔõ $&¤9¨TÔ-?Ô-LÈkÐghÐ#iÑ#iÔ#iˆÔ õ $&¤9¨TÔ-?Ô-LÈaÐ]^Ð#_Ñ#_Ô#_ˆÔ å!.¨vÑ!6Ô!6ˆÔÝ&(¤l°;Ñ&?Ô&?ˆÔ#Ý)+¬°F´NÑ)CÔ)CˆÔ&à�ŠÑÔÐÐÐr4   NÚdecoder_queriesr·  rR   Úprompt_featuresrè  r±   r[   c                 óÀ  — |�`|}|                       |¦  «        }d}	|�t          | j        |||¬¦  «        }	 | j        d||||	dœ|¤Ž\  }
}||                      |
¦  «        z   }|                      ||¬¦  «        }|                      |¦  «        }|                      |¦  «        }t          j	        d||¦  «        }|  
                    |¦  «        }t          ||¬¦  «        S )aZ  
        Args:
            decoder_queries: Decoder output queries [batch_size, num_queries, hidden_size]
            backbone_features: List of backbone features to process through FPN
            encoder_hidden_states: Encoder outputs [batch_size, seq_len, hidden_size]
            prompt_features: Prompt features (text + geometry) for cross-attention [batch_size, prompt_len, hidden_size]
            prompt_mask: Padding mask [batch_size, prompt_len] where True=valid, False=padding

        Returns:
            Sam3MaskDecoderOutput containing predicted masks and semantic segmentation.
        N)r˜   r  rR   r9   rê  )r·  rR   zbqc,bchw->bqhw)rJ   rK   r3   )rÇ  r   r˜   rÆ  rÈ  Ú_embed_pixelsrÄ  rÃ  r0   ÚeinsumrÅ  rI   )r¤   rÉ  r·  rR   rÊ  rè  r±   r`  Únormed_hidden_statesÚcross_attn_maskr¼   r  Úpixel_embedÚinstance_embedsÚmask_embeddingsrJ   rK   s                    r5   r§   zSam3MaskDecoder.forward  s6  € ð, Ð&à,ˆHØ#'×#>Ò#>Ð?TÑ#UÔ#UÐ à"ˆOØÐ&Ý";Øœ;Ø"6Ø*9Ø#.ð	#ñ #ô #�ð 4˜TÔ3ð Ø*Ø#Ø%Ø.ð	ð ð
 ðð ‰NˆK˜ð %-¨t×/MÒ/MÈkÑ/ZÔ/ZÑ$ZÐ!ð ×(Ò(Ø/Ø"7ð )ñ 
ô 
ˆð ×2Ò2°;Ñ?Ô?ˆØ×,Ò,¨_Ñ=Ô=ˆÝ”\Ð"2°OÀ_ÑUÔUˆ
ð ×/Ò/°Ñ<Ô<ˆå$Ø!Ø%ð
ñ 
ô 
ð 	
r4   c                 ó`  — d„ |D ¦   «         }|d         j         d         |d         j         d         z  }|dd…d|…dd…f         }|j         \  }}}|d         j         dd…         \  }	}
|                     dd¦  «                             |||	|
¦  «        }||d<   |                      |¦  «        }|S )aº  
        Embed pixels by combining backbone FPN features with encoder vision features.
        The encoder vision features replace the finest-resolution backbone feature.

        Args:
            backbone_features: List of backbone features [batch_size, C, H_i, W_i]
            encoder_hidden_states: Encoder outputs [batch_size, seq_len, hidden_size]

        Returns:
            Pixel embeddings [batch_size, hidden_size, H, W]
        c                 ó6   — g | ]}|                      ¦   «         ‘ŒS r3   )Úclone)r~  Úfeats     r5   r€  z1Sam3MaskDecoder._embed_pixels.<locals>.<listcomp>[  s    € Ð LÐ LÐ L°$ §¢¡¤Ð LÐ LÐ Lr4   rh   rü   Nr   r   )rn   r·   rÒ   rÂ  )r¤   r·  rR   Úbackbone_visual_featsÚspatial_dimÚencoder_visual_embedrz   r  r|   r  r  rÐ  s               r5   rÌ  zSam3MaskDecoder._embed_pixelsK  sß   € ð  !MÐ LÐ:KÐ LÑ LÔ LÐð (¨Ô+Ô1°"Ô5Ð8IÈ"Ô8MÔ8SÐTVÔ8WÑWˆØ4°Q°Q°Q¸¸¸ÀaÀaÀaÐ5GÔHÐØ%9Ô%?Ñ"ˆ
�A�{Ø)¨"Ô-Ô3°B°C°CÔ8‰ˆ�Ø3×=Ò=¸aÀÑCÔC×KÒKÈJÐXcÐekÐmrÑsÔsÐð %9Ð˜bÑ!ð ×(Ò(Ð)>Ñ?Ô?ˆàÐr4   r]  )r+   r,   r-   r.   r¿   rŒ  r$   r—   r   r   r0   r   rT  r   r   r/   rI   r§   rÌ  r©   rª   s   @r5   rÀ  rÀ  ê  s;  ø€ € € € € ðð ð 	�mðÐðÐ4ð ð ð ð ð ð ð.  Øð 04Ø+/ð<
ð <
àœð<
ð   ¤Ô-ð<
ð  %œ|ð	<
ð
 œ¨Ñ,ð<
ð ”\ DÑ(ð<
ð Ð+Ô,ð<
ð 
Ð&Ñ	&ð<
ð <
ð <
ñ „_ñ  Ôð<
ð|à ¤Ô-ðð  %œ|ðð 
Œð	ð ð ð ð ð ð ð r4   rÀ  c                   ó¨  ‡ — e Zd ZddgZdZddgZdefˆ fd„Zee		 dd	e
j        d
e
j        dz  dee         deez  fd„¦   «         ¦   «         Ze	de
j        dee         defd„¦   «         Zee		 	 	 	 	 	 	 dde
j        dz  dedz  d	e
j        dz  d
e
j        dz  de
j        dz  de
j        dz  de
j        dz  dee         defd„¦   «         ¦   «         Zˆ xZS )Ú	Sam3Modelrf  rg  Údetector_modelz^tracker_model.z^tracker_neck.r˜   c                 óN  •— t          |d¦  «        r1|j        �*|j        }t          |t          ¦  «        rt	          di |¤Ž}|}t          ¦   «                              |¦  «         t          |j        ¦  «        | _	        t          |j        ¦  «        | _        |j        j        | _        t          j        |j        j        |j        j        ¦  «        | _        |j        |j        _        |j        |j        _        |j        |j        _        |j        |j        _        t/          |j        ¦  «        | _        t3          |j        ¦  «        | _        t7          |j        ¦  «        | _        t;          |j        ¦  «        | _        t?          |¦  «        | _         |  !                    ¦   «          d S )NÚdetector_configr3   )"ÚhasattrrÞ  r  Údictr    r–   r—   rÖ  Úvision_configÚvision_encoderr	   Útext_configÚtext_encoderÚ
vocab_sizer›   rœ   r|   Údetr_encoder_configÚtext_projectionrÎ   Úgeometry_encoder_configr�  Úmask_decoder_configrî  Úgeometry_encoderr,  Údetr_encoderr_  Údetr_decoderrÀ  Úmask_decoderr�  Údot_product_scoringr†  )r¤   r˜   rÞ  r¥   s      €r5   r—   zSam3Model.__init__u  sm  ø€ å�6Ð,Ñ-Ô-ð 	%°&Ô2HÐ2TØ$Ô4ˆOÝ˜/­4Ñ0Ô0ð @Ý",Ð"?Ð"?¨Ð"?Ð"?�Ø$ˆFÝ‰Œ×Ò˜Ñ Ô Ð Ý-¨fÔ.BÑCÔCˆÔÝ7¸Ô8JÑKÔKˆÔØ Ô,Ô7ˆŒõ  "œy¨Ô);Ô)GÈÔIcÔIoÑpÔpˆÔð ?EÔ>YˆÔ&Ô;Ø:@Ô:UˆÔ"Ô7Ø:@Ô:UˆÔ"Ô7Ø:@Ô:UˆÔ"Ô7å 3°FÔ4RÑ SÔ SˆÔÝ+¨FÔ,FÑGÔGˆÔÝ+¨FÔ,FÑGÔGˆÔÝ+¨FÔ,FÑGÔGˆÔõ $9¸Ñ#@Ô#@ˆÔ à�ŠÑÔÐÐÐr4   NÚ	input_idsr9   r±   r[   c                 ój   —  | j         d||ddœ|¤Ž}|j        }|                      |¦  «        |_        |S )a´  
        Example:

        ```python
        >>> from transformers import Sam3Model, Sam3Processor
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO

        >>> model = Sam3Model.from_pretrained("facebook/sam3")
        >>> processor = Sam3Processor.from_pretrained("facebook/sam3")

        >>> # Pre-compute text embeddings
        >>> text_inputs = processor(text="cat", return_tensors="pt")
        >>> text_embeds = model.get_text_features(**text_inputs).pooler_output

        >>> # Reuse text embeddings for multiple images
        >>> url = "http://images.cocodataset.org/val2017/000000077595.jpg"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read()))
        >>> img_inputs = processor(images=image, return_tensors="pt")
        >>> outputs = model(pixel_values=img_inputs.pixel_values, text_embeds=text_embeds)
        ```
        T©rï  r9   Úreturn_dictr3   )rä  r8   rç  Úpooler_output)r¤   rï  r9   r±   Útext_outputsr8   s         r5   Úget_text_featureszSam3Model.get_text_features•  sZ   € ð@ )�tÔ(ð 
Ø°ÈDð
ð 
ØTZð
ð 
ˆð )Ô:ÐØ%)×%9Ò%9Ð:KÑ%LÔ%LˆÔ"àÐr4   r   c                 ó"   —  | j         |fi |¤Ž}|S )aÊ  
        Example:

        ```python
        >>> from transformers import Sam3Model, Sam3Processor
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO

        >>> model = Sam3Model.from_pretrained("facebook/sam3")
        >>> processor = Sam3Processor.from_pretrained("facebook/sam3")

        >>> # Pre-compute vision embeddings
        >>> url = "http://images.cocodataset.org/val2017/000000077595.jpg"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read()))
        >>> img_inputs = processor(images=image, return_tensors="pt")
        >>> vision_embeds = model.get_vision_features(pixel_values=img_inputs.pixel_values)

        >>> # Reuse vision embeddings for multiple text prompts
        >>> text_inputs = processor(text="cat", return_tensors="pt")
        >>> outputs = model(vision_embeds=vision_embeds, input_ids=text_inputs.input_ids)
        ```
        )râ  )r¤   r   r±   Úvision_outputss       r5   Úget_vision_featureszSam3Model.get_vision_features½  s$   € ð< -˜Ô,¨\ÐDÐD¸VÐDÐDˆØÐr4   Úvision_embedsÚtext_embedsÚinput_boxesÚinput_boxes_labelsc                 óÖ	  — |du |du k    rt          d¦  «        ‚|du |du k    rt          d¦  «        ‚|�|j        d         }	|j        }
n*|j        d         j        d         }	|j        d         j        }
|€ | j        |fi |¤Ž}n|}|j        dd…         }|j        dd…         }|€|                      ||d¬¦  «        }|j        }|�|                     ¦   «         nd}|duo| 	                    ¦   «         dk    }d}d}|�r |�”| 	                    ¦   «         dk    r||}|�|n%t          j        |d         t          j        ¬	¦  «        }|�|d
k    n,t          j        |	|j        d         t          j        |
¬¦  «        }t          j        |d
k    d|¦  «        }nbt          j        |	dd|j        |
¬¦  «        }t          j        |	dt          j        |
¬¦  «        }t          j        |	dt          j        |
¬¦  «        }|                      |||||¬¦  «        }|j        }|j        }|��Q|j        d         dk    r3|j        d         dk    r"|                     |j        d         dd¦  «        }t          j        ||gd¬¦  «        }|�C|j        d         dk    r2|j        d         dk    r!|                     |j        d         d¦  «        }|�|�t          j        ||gd¬¦  «        }n—|�Ft          j        |	|j        d         t          j        |
¬¦  «        }t          j        ||gd¬¦  «        }nO|�Ft          j        |	|j        d         t          j        |
¬¦  «        }t          j        ||gd¬¦  «        }nd}n|}|} | j        d|d         g||d         g|dœ|¤Ž} | j        d|j        |j        |j        ||j        dœ|¤Ž}| j                             |j        ¦  «        }t;          |j        ¦  «        }||z                        ¦   «         }tA          |¦  «        } |  !                    |j        |j        |¬¦  «         "                    d¦  «        }!|!d         }"| d         }#|j        d         }$|j#        d         }% | j$        d|$tK          |¦  «        |j        ||dœ|¤Ž}&tM          |&j'        |#|"|%|&j(        |j)        |j        |j)        |j)        |j*        |j*        |j*        |&j*        ¬¦  «        S )aÞ  
        vision_embeds (`Sam3VisionEncoderOutput`, *optional*):
            Pre-computed vision embeddings. Can be used to easily reuse vision embeddings. If provided, `pixel_values`
            should not be passed. Mutually exclusive with `pixel_values`.
        text_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
            Pre-computed text embeddings. Can be used to easily reuse text embeddings. If provided, `input_ids`
            should not be passed. Mutually exclusive with `input_ids`.
        input_boxes (`torch.FloatTensor` of shape `(batch_size, num_boxes, 4)`, *optional*):
            Normalized box coordinates in [0, 1] range, in (cx, cy, w, h) format.
        input_boxes_labels (`torch.LongTensor` of shape `(batch_size, num_boxes)`, *optional*):
            Labels for boxes: 1 (positive), 0 (negative).

        Example:

        ```python
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO
        >>> from transformers import AutoModel, AutoProcessor

        >>> model = AutoModel.from_pretrained("facebook/sam3")
        >>> processor = AutoProcessor.from_pretrained("facebook/sam3")

        >>> url = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/transformers/model_doc/sam-car.png"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read())).convert("RGB")
        >>> text = "car"
        >>> inputs = processor(images=image, text=text, return_tensors="pt")

        >>> # Get segmentation output
        >>> outputs = model(**inputs)
        >>> pred_masks = outputs.pred_masks
        >>> pred_boxes = outputs.pred_boxes
        ```
        Nz=You must specify exactly one of pixel_values or vision_embedsz8You must specify exactly one of input_ids or text_embedsr   rh   Trñ  ).r   rÞ   iöÿÿÿr   r™  rÝ   )r  r  r  r  r  ri   )r
  r>   r  r5  )r
  r>   rç  r5  r?   )rP   r>   r5  )rÉ  r·  rR   rÊ  rè  )rJ   rN   rO   rG   rK   rP   rQ   rR   rS   rT   rU   rV   rW   r3   )+rå   rn   rk   r)   râ  r*   rõ  ró  rÏ   Únumelr0   Ú	ones_likerè   rM  Úwherers   rm   rê  r8   r9   rr   rë   rë  rì  r>   r=   r?   rb  rE   re   rF   r~  r‘   rî  r€  rG   rí  rT  rM   rJ   rK   r@   rA   )'r¤   r   rù  rï  r9   rú  rû  rü  r±   rz   rk   r÷  r)   r*   r>   r5  Úhas_geometry_promptsÚgeometry_prompt_featuresÚgeometry_prompt_maskr  r  r  Úgeometry_outputsÚcombined_prompt_featuresÚcombined_prompt_maskÚgeo_valid_maskÚtext_valid_maskÚencoder_outputsÚdecoder_outputsÚall_box_offsetsÚreference_boxes_inv_sigÚall_pred_boxes_cxcywhÚall_pred_boxesÚall_pred_logitsrO   rN   rP   rG   Úmask_outputss'                                          r5   r§   zSam3Model.forwardÞ  s  € ð` ˜DÐ  m°tÐ&;Ò<Ð<ÝÐ\Ñ]Ô]Ð]à˜Ð ;°$Ð#6Ò7Ð7ÝÐWÑXÔXÐXàÐ#Ø%Ô+¨AÔ.ˆJØ!Ô(ˆFˆFà&Ô8¸Ô;ÔAÀ!ÔDˆJØ"Ô4°QÔ7Ô>ˆFàÐ Ø0˜TÔ0°ÐHÐHÀÐHÐHˆNˆNà*ˆNà*Ô<¸S¸b¸SÔAÐØ .Ô DÀSÀbÀSÔ IÐàÐØ×0Ò0¸9ÐUcÐquÐ0ÑvÔvˆKà#Ô1ˆØ-;Ð-G�N×'Ò'Ñ)Ô)Ð)ÈTˆ	Ø*°$Ð6ÐR¸;×;LÒ;LÑ;NÔ;NÐQRÒ;RÐà#'Ð Ø#Ðàñ 	CØÐ&¨;×+<Ò+<Ñ+>Ô+>ÀÒ+BÐ+BØ!,�ð *Ð5ð 'Ð&åœ¨¸Ô)?ÅuÄzÐRÑRÔRð ð *Ð5ð (¨3Ò.Ð.åœ J°Ô0AÀ!Ô0DÍEÌJÐ_eÐfÑfÔfð õ
 #œ[¨°sÒ):¸A¸zÑJÔJ�
�
å!&¤¨Z¸¸AÀ]ÔEXÐagÐ!hÑ!hÔ!h�Ý"œ[¨°Q½e¼jÐQWÐXÑXÔX�
Ý œ; z°1½E¼JÈvÐVÑVÔV�à#×4Ò4Ø-Ø!Ø%Ø+Ø4ð  5ñ  ô  Ðð (8Ô'IÐ$Ø#3Ô#BÐ à#Ñ/àÔ" 1Ô%¨Ò*Ð*Ð/GÔ/MÈaÔ/PÐSTÒ/TÐ/TØ -× 4Ò 4Ð5MÔ5SÐTUÔ5VÐXYÐ[\Ñ ]Ô ]�Ý',¤y°-ÐAYÐ1ZÐ`aÐ'bÑ'bÔ'bÐ$ØÐ$¨¬¸Ô);¸qÒ)@Ð)@ÐEYÔE_Ð`aÔEbÐefÒEfÐEfØ%×,Ò,Ð-AÔ-GÈÔ-JÈAÑNÔN�	àÐ$Ð)=Ð)IÝ',¤y°)Ð=QÐ1RÐXYÐ'ZÑ'ZÔ'ZÐ$Ð$ØÐ&Ý!&¤ØÐ 8Ô >¸qÔ AÍÌÐ\bð"ñ "ô "�õ (-¤y°)¸^Ð1LÐRSÐ'TÑ'TÔ'TÐ$Ð$Ø%Ð1Ý"'¤*¨Z¸Ô9LÈQÔ9OÕW\ÔWaÐjpÐ"qÑ"qÔ"q�Ý',¤y°/ÐCWÐ1XÐ^_Ð'`Ñ'`Ô'`Ð$Ð$à'+Ð$Ð$à'4Ð$Ø#,Ð à+˜$Ô+ð 
Ø.¨rÔ2Ð3Ø2Ø4°RÔ8Ð9Ø*ð	
ð 
ð
 ð
ð 
ˆð ,˜$Ô+ð 
Ø+Ô=Ø)Ô7Ø /Ô DØ*Ø*Ô9ð
ð 
ð ð
ð 
ˆð Ô+×4Ò4°_Ô5_Ñ`Ô`ˆÝ"1°/Ô2QÑ"RÔ"RÐØ!8¸?Ñ!J× SÒ SÑ UÔ UÐÝ+Ð,AÑBÔBˆà×2Ò2Ø"1Ô"LØ)Ô7Ø*ð 3ñ 
ô 
÷ Š'�"‰+Œ+ð	 	ð & bÔ)ˆØ# BÔ'ˆ
Ø /Ô JÈ2Ô NÐØ)Ô9¸"Ô=ˆà(�tÔ(ð 
Ø1Ý"Ð#4Ñ5Ô5Ø"1Ô"CØ4Ø,ð
ð 
ð ð
ð 
ˆõ +Ø#Ô.Ø!Ø#Ø+Ø%Ô2Ø"1Ô"?Ø$3Ô$CØ"1Ô"?Ø!/Ô!=Ø,Ô7Ø$3Ô$>Ø$3Ô$>Ø$0Ô$;ð
ñ 
ô 
ð 	
r4   r•   )NNNNNNN)r+   r,   r-   rs  rq  Ú"_keys_to_ignore_on_load_unexpectedr    r—   r   r   r0   rB   r   r   r   r/   r   rõ  r1   r(   rø  rM   r§   r©   rª   s   @r5   rÛ  rÛ  m  s  ø€ € € € € Ø Ð(ÐØ(ÐàØð*Ð&ð
˜zð ð ð ð ð ð ð@ Øð /3ð$ð $àÔ#ð$ð œ tÑ+ð$ð Ð+Ô,ð	$ð
 
Ð+Ñ	+ð$ð $ð $ñ „^ñ Ôð$ðL ðàÔ'ðð Ð+Ô,ðð 
!ð	ð ð ñ „^ðð@ Øð 26Ø8<Ø-1Ø.2Ø04Ø04Ø6:ð|
ð |
àÔ'¨$Ñ.ð|
ð /°Ñ5ð|
ð Ô# dÑ*ð	|
ð
 œ tÑ+ð|
ð Ô&¨Ñ-ð|
ð Ô&¨Ñ-ð|
ð "Ô,¨tÑ3ð|
ð Ð+Ô,ð|
ð 
%ð|
ð |
ð |
ñ „^ñ Ôð|
ð |
ð |
ð |
ð |
r4   rÛ  )rÛ  rÖ  ry  rd  )rX   r9  )Nr«   )mr•  Úcollections.abcr   r   Údataclassesr   Únumpyr•  r0   Útorch.nnr›   Útorch.nn.functionalr¸   rD  r   Úutilsr   ra  Útransformersr	   Ú r
   rl  Úactivationsr   Úmasking_utilsr   Úmodeling_layersr   Úmodeling_outputsr   r   r   Úmodeling_utilsr   r   Úprocessing_utilsr   Úpytorch_utilsr   r   r   r   Úutils.genericr   r   r   Úutils.import_utilsr   Úutils.output_capturingr   Úautor   Úconfiguration_sam3r    r!   r"   r#   r$   r%   r&   Ú
get_loggerr+   rÐ   r(   r7   r<   rD   rI   rM   rç   re   rÏ   r‡   r‘   ÚModuler“   r½   r¿   r×   rÿ   r/   r  r  r  r&  rD  rH  rJ  rQ  rd  ry  rŽ  rµ  rÇ  rÖ  rà  rî  r&  r,  r;  rG  r_  r�  r¥  r®  rÀ  rÛ  Ú__all__r3   r4   r5   ú<module>r)     sÛ
  ðð €€€Ø .Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ø !Ð !Ð !Ð !Ð !Ð !à Ð Ð Ð Ø €€€Ø Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð à -Ð -Ð -Ð -Ð -Ð -ð ÐÑÔð ØÐÐÐà 4Ð 4Ð 4Ð 4Ð 4Ð 4à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø 6Ð 6Ð 6Ð 6Ð 6Ð 6Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9ðð ð ð ð ð ð ð ð ð ð
 GÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø @Ð @Ð @Ð @Ð @Ð @Ø >Ð >Ð >Ð >Ð >Ð >Ð >Ð >Ð >Ð >ðð ð ð ð ð ð ð ð ð ð
 +Ð *Ð *Ð *Ð *Ð *Ø 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø Ð Ð Ð Ð Ð ðð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð 
ˆÔ	˜HÑ	%Ô	%€ð Ø
ð	@ð 	@ð 	@ð 	@ð 	@Ð8ñ 	@ô 	@ñ „ñ „ð	@ð Ø
ð	3ð 	3ð 	3ð 	3ð 	3 ñ 	3ô 	3ñ „ñ „ð	3ð Ø
ð7ð 7ð 7ð 7ð 7˜Kñ 7ô 7ñ „ñ „ð7ð0 Ø
ð7ð 7ð 7ð 7ð 7˜Kñ 7ô 7ñ „ñ „ð7ð* Ø
ð7ð 7ð 7ð 7ð 7˜Kñ 7ô 7ñ „ñ „ð7ð Ø
ð-Dð -Dð -Dð -Dð -D +ñ -Dô -Dñ „ñ „ð-Dð`ð �u”|ð ¨%ð ¸5¼<ð ð ð ð ð34ð 34ÀDð 34ð 34ð 34ð 34ðl"ð "ð "ðð ð ð ð ˆbŒiñ ô ð ð. !Øð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð �T‰\ð%ð ð%ð Ð'Ô(ð%ð %ð %ð %ð8P)ð P)ð P)ð P)ð P)�B”Iñ P)ô P)ð P)ðf Bð  Bð  Bð  Bð  B˜RœYñ  Bô  Bð  BðF#ð #ð #ð$2Ø„|ð2à„|ð2ð 
Œð2ð 
Œð	2ð
 ˆ5Œ<˜œÐ%Ô&ð2ð 2ð 2ð 2ð62)ð 2)ð 2)ð 2)ð 2)˜2œ9ñ 2)ô 2)ð 2)ðjð ð ð ð ˜RœYñ ô ð ð6Eð Eð Eð Eð E˜œ	ñ Eô Eð EðP2ð 2ð 2ð>ð ð ð>+ð +ð +ð +ð +˜œ	ñ +ô +ð +ð6ð 6ð 6ð 6ð 6Ð-ñ 6ô 6ð 6ðr Ø	€Ð+Ð,Ñ,Ô,ðCð Cð Cð Cð C˜/ñ Cô Cñ -Ô,ñ „ðCð: ð0@ð 0@ð 0@ð 0@ð 0@Ð&ñ 0@ô 0@ñ „ð0@ðfB
ð B
ð B
ð B
ð B
 ¤	ñ B
ô B
ð B
ðJ#ð #ð #ð #ð #�2”9ñ #ô #ð #ðL8ð 8ð 8ð 8ð 8�R”Yñ 8ô 8ð 8ðB €ððñ ô ð
(
ð (
ð (
ð (
ð (
Ð)ñ (
ô (
ñô ð
(
ðV%ð %ð %ð %ð %˜rœyñ %ô %ð %ðPc
ð c
ð c
ð c
ð c
˜"œ)ñ c
ô c
ð c
ðLCð Cð Cð Cð C˜2œ9ñ Cô Cð CðL
ð 
ð 
ð 
ð 
Ð)ñ 
ô 
ð 
ðDð ð ð ð �R”Yñ ô ð ð4bð bð bð bð b˜2œ9ñ bô bð bðJN
ð N
ð N
ð N
ð N
Ð)ñ N
ô N
ð N
ðb[ð [ð [ð [ð [˜BœIñ [ô [ð [ð|!ð !ð !ð !ð !�r”yñ !ô !ð !ðH/ð /ð /ð /ð /�r”yñ /ô /ð /ðd@ð @ð @ð @ð @Ð)ñ @ô @ð @ðFo
ð o
ð o
ð o
ð o
Ð#ñ o
ô o
ð o
ðd	 RÐ
QÐ
Q€€€r4   