§
    ‚Štj(Ô  ã                   ó0  — d dl Zd dlZd dl mZ d dlmZ d dlZd dlZd dl	m
c 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 dd
lmZmZ ddlmZ ddlmZm Z m!Z! ddl"m#Z# ddl$m%Z% ddl&m'Z'  e¦   «         rd dl(m)Z)  e!¦   «         rd dl*m+Z+ d dl,m-Z-  e d¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z.	 dNdej        dej        dej        fd„Z/dededefd„Z0dej        dej        dej        fd „Z1 G d!„ d"e
j2        ¦  «        Z3deded#e4defd$„Z5dej        dej        d#e4dej        fd%„Z6 G d&„ d'e
j2        ¦  «        Z7 G d(„ d)e
j2        ¦  «        Z8 G d*„ d+e
j2        ¦  «        Z9	 dOd-e
j2        d.ej        d/ej        d0ej        d1ej        dz  d2e:d3e:fd4„Z; G d5„ d6e
j2        ¦  «        Z< G d7„ d8e
j2        ¦  «        Z= G d9„ d:e
j2        ¦  «        Z> G d;„ d<e
j2        ¦  «        Z? G d=„ d>e
j2        ¦  «        Z@ G d?„ d@e¦  «        ZA G dA„ dBe
jB        ¦  «        ZC G dC„ dDe
j2        ¦  «        ZD G dE„ dFe
j2        ¦  «        ZE G dG„ dHe
j2        ¦  «        ZFe  G dI„ dJe¦  «        ¦   «         ZG e dK¬¦  «         G dL„ dMeG¦  «        ¦   «         ZHdJdMgZIdS )Pé    N)ÚCallable)Ú	dataclass)ÚTensorÚnné   )Úinitialization)ÚACT2FN)ÚModelOutputÚis_scipy_availableÚrequires_backends)ÚGradientCheckpointingLayer)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)ÚTransformersKwargsÚauto_docstringÚis_accelerate_available)Úmerge_with_config_defaults)Úcapture_outputsé   )Ú
EomtConfig)Úlinear_sum_assignment)ÚPartialState)Úreducea˜  
    Class for outputs of [`EomtForUniversalSegmentationOutput`].

    This output can be directly passed to [`~EomtImageProcessor.post_process_semantic_segmentation`] or
    [`~EomtImageProcessor.post_process_instance_segmentation`] or
    [`~EomtImageProcessor.post_process_panoptic_segmentation`] to compute final segmentation maps. Please, see
    [`~EomtImageProcessor] for details regarding usage.
    )Úcustom_introc                   ó
  — e Zd ZU dZ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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 )
Ú"EomtForUniversalSegmentationOutputa*  
    loss (`torch.Tensor`, *optional*):
        The computed loss, returned when labels are present.
    class_queries_logits (`torch.FloatTensor`):
        A tensor of shape `(batch_size, num_queries, num_labels + 1)` representing the proposed classes for each
        query. Note the `+ 1` is needed because we incorporate the null class.
    masks_queries_logits (`torch.FloatTensor`):
        A tensor of shape `(batch_size, num_queries, height, width)` representing the proposed masks for each
        query.
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
        Last hidden states (final feature map) of the last layer.
    hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each stage) of
        shape `(batch_size, sequence_length, hidden_size)`. Hidden-states all layers of the model.
    attentions (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `tuple(torch.FloatTensor)` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`. Self and Cross Attentions weights from transformer decoder.
    patch_offsets (`list[torch.Tensor]`, *optional*):
        list of tuples indicating the image index and start and end positions of patches for semantic segmentation.
    NÚlossÚclass_queries_logitsÚmasks_queries_logitsÚlast_hidden_stateÚhidden_statesÚ
attentionsÚpatch_offsets)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚtorchÚFloatTensorÚ__annotations__r   r    r!   r"   Útupler#   r$   Úlistr   © ó    úd/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/eomt/modeling_eomt.pyr   r   3   s×   € € € € € € ðð ð* &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø59Ð˜%Ô+¨dÑ2Ð9Ð9Ñ9Ø59Ð˜%Ô+¨dÑ2Ð9Ð9Ñ9Ø26Ð�uÔ(¨4Ñ/Ð6Ð6Ñ6Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ø26€J��eÔ'Ô(¨4Ñ/Ð6Ð6Ñ6Ø/3€M�4˜œÔ%¨Ñ,Ð3Ð3Ñ3Ð3Ð3r/   r   FÚinput_featuresÚpoint_coordinatesÚreturnc                 óØ   — |                      ¦   «         dk    rd}|                     d¦  «        }t          j        j        j        | d|z  dz
  fi |¤Ž}|r|                     d¦  «        }|S )a(  
    A wrapper around `torch.nn.functional.grid_sample` to support 3D point_coordinates tensors.

    Args:
        input_features (`torch.Tensor` of shape (batch_size, channels, height, width)):
            A tensor that contains features map on a height * width grid
        point_coordinates (`torch.Tensor` of shape (batch_size, num_points, 2) or (batch_size, grid_height, grid_width,:
        2)):
            A tensor that contains [0, 1] * [0, 1] normalized point coordinates
        add_dim (`bool`):
            boolean value to keep track of added dimension

    Returns:
        point_features (`torch.Tensor` of shape (batch_size, channels, num_points) or (batch_size, channels,
        height_grid, width_grid):
            A tensor that contains features for points in `point_coordinates`.
    r   Té   g       @ç      ð?)ÚdimÚ	unsqueezer)   r   Ú
functionalÚgrid_sampleÚsqueeze)r1   r2   Úadd_dimÚkwargsÚpoint_featuress        r0   Úsample_pointr?   ^   sƒ   € ð( ×ÒÑÔ !Ò#Ð#ØˆØ-×7Ò7¸Ñ:Ô:Ðõ ”XÔ(Ô4°^ÀSÐK\ÑE\Ð_bÑEbÐmÐmÐflÐmÐm€NØð 3Ø'×/Ò/°Ñ2Ô2ˆàÐr/   ÚinputsÚlabelsc                 ó(  — |                       ¦   «                              d¦  «        } dt          j        | |j        ¦  «        z  }|                      d¦  «        dd…df         |                     d¦  «        ddd…f         z   }d|dz   |dz   z  z
  }|S )aÈ  
    A pair wise version of the dice loss, see `dice_loss` for usage.

    Args:
        inputs (`torch.Tensor`):
            A tensor representing a mask
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).

    Returns:
        `torch.Tensor`: The computed loss between each pairs.
    r   r5   éÿÿÿÿN)ÚsigmoidÚflattenr)   ÚmatmulÚTÚsum)r@   rA   Ú	numeratorÚdenominatorr   s        r0   Úpair_wise_dice_lossrK   ~   s�   € ð �^Š^ÑÔ×%Ò% aÑ(Ô(€FØ•E”L ¨¬Ñ2Ô2Ñ2€Ià—*’*˜R‘.”.    D Ô)¨F¯JªJ°r©N¬N¸4ÀÀÀ¸7Ô,CÑC€KØ�	˜A‘ +°¡/Ñ2Ñ2€DØ€Kr/   c                 óF  — | j         d         }t          j        d¬¦  «        } || t          j        | ¦  «        ¦  «        } || t          j        | ¦  «        ¦  «        }t          j        ||z  |j        ¦  «        }t          j        ||z  d|z
  j        ¦  «        }||z   }|S )aê  
    A pair wise version of the cross entropy loss, see `sigmoid_cross_entropy_loss` for usage.

    Args:
        inputs (`torch.Tensor`):
            A tensor representing a mask.
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).

    Returns:
        loss (`torch.Tensor`): The computed loss between each pairs.
    r   Únone©Ú	reduction)Úshaper   ÚBCEWithLogitsLossr)   Ú	ones_likeÚ
zeros_likerF   rG   )	r@   rA   Úheight_and_widthÚ	criterionÚcross_entropy_loss_posÚcross_entropy_loss_negÚloss_posÚloss_negr   s	            r0   Ú$pair_wise_sigmoid_cross_entropy_lossrZ   ”   s£   € ð ”| A”ÐåÔ$¨vÐ6Ñ6Ô6€IØ&˜Y v­u¬¸vÑ/FÔ/FÑGÔGÐØ&˜Y v­uÔ/?ÀÑ/GÔ/GÑHÔHÐåŒ|Ð2Ð5EÑEÀvÄxÑPÔP€HÝŒ|Ð2Ð5EÑEÈÈFÉ
Ä~ÑVÔV€HØ�hÑ€DØ€Kr/   c                   ó¾   ‡ — e Zd ZdZ	 ddedededefˆ fd„Z ej        ¦   «         d	ej	        d
ej	        dej	        dej	        de
ee	                  f
d„¦   «         Zˆ xZS )ÚEomtHungarianMatcheraq  This class computes an assignment between the labels and the predictions of the network.

    For efficiency reasons, the labels don't include the no_object. Because of this, in general, there are more
    predictions than labels. In this case, we do a 1-to-1 matching of the best predictions, while the others are
    un-matched (and thus treated as non-objects).
    r6   é 1  Ú
cost_classÚ	cost_maskÚ	cost_diceÚ
num_pointsc                 óÂ   •— t          ¦   «                              ¦   «          |dk    r|dk    r|dk    rt          d¦  «        ‚|| _        || _        || _        || _        dS )aH  Creates the matcher

        Params:
            cost_class (`float`, *optional*, defaults to 1.0):
                Relative weight of the classification error in the matching cost.
            cost_mask (`float`, *optional*,  defaults to 1.0):
                This is the relative weight of the focal loss of the binary mask in the matching cost.
            cost_dice (`float`, *optional*, defaults to 1.0):
                This is the relative weight of the dice loss of the binary mask in the matching cost.
            num_points (`int`, *optional*, defaults to 12544):
                No. of points to sample on which the mask loss will be calculated. The same set of K points are
                uniformly sampled for all prediction and ground truth masks to construct the cost matrix for bipartite
                matching.
        r   zAll costs can't be 0N)ÚsuperÚ__init__Ú
ValueErrorra   r^   r_   r`   )Úselfr^   r_   r`   ra   Ú	__class__s        €r0   rd   zEomtHungarianMatcher.__init__¸   sc   ø€ õ" 	‰Œ×ÒÑÔÐØ˜Š?ˆ?˜y¨Aš~˜~°)¸q².°.ÝÐ3Ñ4Ô4Ð4à$ˆŒØ$ˆŒØ"ˆŒØ"ˆŒˆˆr/   r    r   Úmask_labelsÚclass_labelsr3   c                 óH  — g }|j         d         }t          |¦  «        D �]õ}||                              d¦  «        }||         }	|dd…||         f          }
||                              |	¦  «        }|dd…df         }|	dd…df         }	t	          j        d| j        d|	j        ¬¦  «        }|                     |j         d         dd¦  «        }t          ||d¬¦  «         
                    d¦  «        }|                     |	j         d         dd¦  «        }t          |	|d¬¦  «         
                    d¦  «        }	t          |	|¦  «        }t          |	|¦  «        }| j        |z  | j        |
z  z   | j        |z  z   }t	          j        |t	          j        d	¦  «        ¦  «        }t	          j        |t	          j        d
¦  «        ¦  «        }t	          j        |d¦  «        }t)          |                     ¦   «         ¦  «        }|                     |¦  «         �Œ÷d„ |D ¦   «         }|S )ao  
        Params:
            masks_queries_logits (`torch.Tensor`):
                A tensor of dim `batch_size, num_queries, num_labels` with the classification logits.
            class_queries_logits (`torch.Tensor`):
                A tensor of dim `batch_size, num_queries, height, width` with the predicted masks.
            class_labels (`torch.Tensor`):
                A tensor of dim `num_target_boxes` (where num_target_boxes is the number of ground-truth objects in the
                target) containing the class labels.
            mask_labels (`torch.Tensor`):
                A tensor of dim `num_target_boxes, height, width` containing the target masks.

        Returns:
            matched_indices (`list[tuple[Tensor]]`): A list of size batch_size, containing tuples of (index_i, index_j)
            where:
                - index_i is the indices of the selected predictions (in order)
                - index_j is the indices of the corresponding selected labels (in order)
            For each batch element, it holds:
                len(index_i) = len(index_j) = min(num_queries, num_target_boxes).
        r   rC   Nr   r5   ©ÚdeviceF©Úalign_cornersg    _ Bg    _ Âc                 ó”   — g | ]E\  }}t          j        |t           j        ¬ ¦  «        t          j        |t           j        ¬ ¦  «        f‘ŒFS )©Údtype)r)   Ú	as_tensorÚint64)Ú.0ÚiÚjs      r0   ú
<listcomp>z0EomtHungarianMatcher.forward.<locals>.<listcomp>  sR   € ð 
ð 
ð 
Ù_cÐ_`Ðbc�UŒ_˜Q¥e¤kÐ2Ñ2Ô2µE´OÀAÍUÌ[Ð4YÑ4YÔ4YÐZð
ð 
ð 
r/   )rP   ÚrangeÚsoftmaxÚtor)   Úrandra   rl   Úrepeatr?   r;   rZ   rK   r_   r^   r`   ÚminimumÚtensorÚmaximumÚ
nan_to_numr   ÚcpuÚappend)rf   r    r   rh   ri   ÚindicesÚ
batch_sizeru   Ú
pred_probsÚ	pred_maskr^   Útarget_maskr2   Útarget_coordinatesÚpred_coordinatesr_   r`   Úcost_matrixÚassigned_indicesÚmatched_indicess                       r0   ÚforwardzEomtHungarianMatcher.forwardÒ   s6  € ð8 *,ˆð *Ô/°Ô2ˆ
Ý�zÑ"Ô"ð 	-ñ 	-ˆAØ-¨aÔ0×8Ò8¸Ñ<Ô<ˆJØ,¨QÔ/ˆIð % Q Q Q¨°Q¬Ð%7Ô8Ð8ˆJØ% aœ.×+Ò+¨IÑ6Ô6ˆKØ% a a a¨ gÔ.ˆKØ! ! ! ! T 'Ô*ˆIõ !&¤
¨1¨d¬o¸qÈÔIYÐ ZÑ ZÔ ZÐà!2×!9Ò!9¸+Ô:KÈAÔ:NÐPQÐSTÑ!UÔ!UÐÝ& {Ð4FÐV[Ð\Ñ\Ô\×dÒdÐefÑgÔgˆKà0×7Ò7¸	¼ÈÔ8JÈAÈqÑQÔQÐÝ$ YÐ0@ÐPUÐVÑVÔV×^Ò^Ð_`ÑaÔaˆIõ =¸YÈÑTÔTˆIå+¨I°{ÑCÔCˆIàœ.¨9Ñ4°t´ÈÑ7SÑSÐVZÔVdÐgpÑVpÑpˆKåœ-¨µU´\À$Ñ5GÔ5GÑHÔHˆKÝœ-¨µU´\À%Ñ5HÔ5HÑIÔIˆKÝÔ*¨;¸Ñ:Ô:ˆKå0EÀkÇoÂoÑFWÔFWÑ0XÔ0XÐØ�NŠNÐ+Ñ,Ô,Ð,Ñ,ð
ð 
Øgnð
ñ 
ô 
ˆð Ðr/   )r6   r6   r6   r]   )r%   r&   r'   r(   ÚfloatÚintrd   r)   Úno_gradr   r-   r,   r�   Ú__classcell__©rg   s   @r0   r\   r\   °   sé   ø€ € € € € ðð ð joð#ð #Øð#Ø27ð#ØJOð#Øcfð#ð #ð #ð #ð #ð #ð4 €U„]�_„_ðDà#œlðDð $œlðDð ”\ð	Dð
 ”lðDð 
ˆe�FŒmÔ	ðDð Dð Dñ „_ðDð Dð Dð Dð Dr/   r\   Ú	num_masksc                 ó*  — |                       ¦   «                              d¦  «        }d||z                       d¦  «        z  }|                     d¦  «        |                     d¦  «        z   }d|dz   |dz   z  z
  }|                     ¦   «         |z  }|S )a4  
    Compute the DICE loss, similar to generalized IOU for masks as follows:

    $$ \mathcal{L}_{\text{dice}(x, y) = 1 - \frac{2 * x \cap y }{x \cup y + 1}} $$

    In practice, since `labels` is a binary mask, (only 0s and 1s), dice can be computed as follow

    $$ \mathcal{L}_{\text{dice}(x, y) = 1 - \frac{2 * x * y }{x + y + 1}} $$

    Args:
        inputs (`torch.Tensor`):
            A tensor representing a mask.
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).
        num_masks (`int`):
            The number of masks present in the current batch, used for normalization.

    Returns:
        `torch.Tensor`: The computed loss.
    r   r5   rC   )rD   rE   rH   )r@   rA   r“   ÚprobsrI   rJ   r   s          r0   Ú	dice_lossr–     s‰   € ð, �NŠNÑÔ×$Ò$ QÑ'Ô'€EØ�U˜V‘^×(Ò(¨Ñ,Ô,Ñ,€IØ—)’)˜B‘-”- &§*¢*¨R¡.¤.Ñ0€KØ�	˜A‘ +°¡/Ñ2Ñ2€DØ�8Š8‰:Œ:˜	Ñ!€DØ€Kr/   c                 óœ   — t          j        d¬¦  «        } || |¦  «        }|                     d¦  «                             ¦   «         |z  }|S )a|  
    Args:
        inputs (`torch.Tensor`):
            A float tensor of arbitrary shape.
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).

    Returns:
        loss (`torch.Tensor`): The computed loss.
    rM   rN   r   )r   rQ   ÚmeanrH   )r@   rA   r“   rU   Úcross_entropy_lossr   s         r0   Úsigmoid_cross_entropy_lossrš   8  sR   € õ Ô$¨vÐ6Ñ6Ô6€IØ"˜ 6¨6Ñ2Ô2Ðà×"Ò" 1Ñ%Ô%×)Ò)Ñ+Ô+¨iÑ7€DØ€Kr/   c                   ó~  ‡ — e Zd Zdedeeef         fˆ fd„Zdeee	                  dee	         fd„Z
dee         deeef         fd„Zd	ed
ee         deej                 deeef         fd„Zdej        deej                 deej                 de	deeej        f         f
d„Zd„ Zd„ Zdej        dej        fd„Zdej        de	de	dedej        f
d„Z	 ddej        d	ej        deej                 d
eej                 deeej        f         dz  deeej        f         fd„Zd
ej        dej        dej        fd„Zˆ xZS )ÚEomtLossÚconfigÚweight_dictc                 óÀ  •— t          ¦   «                              ¦   «          t          | dg¦  «         |j        | _        || _        |j        | _        t          j        | j        dz   ¦  «        }| j        |d<   |  	                    d|¦  «         |j
        | _        |j        | _        |j        | _        t          |j        |j        |j        | j        ¬¦  «        | _        dS )aH  
        The Eomt Loss. The loss is computed very similar to DETR. The process happens in two steps: 1) we
        compute hungarian assignment between ground truth masks and the outputs of the model 2) we supervise each pair
        of matched ground-truth / prediction (supervise class and mask)

        Args:
            config (`EomtConfig`):
                The configuration for Eomt model also containing loss calculation specific parameters.
            weight_dict (`dict[str, float]`):
                A dictionary of weights to be applied to the different losses.
        Úscipyr   rC   Úempty_weight)r^   r`   r_   ra   N)rc   rd   r   Ú
num_labelsrž   Úno_object_weightÚeos_coefr)   ÚonesÚregister_bufferÚtrain_num_pointsra   Úoversample_ratioÚimportance_sample_ratior\   Úclass_weightÚdice_weightÚmask_weightÚmatcher)rf   r�   rž   r¡   rg   s       €r0   rd   zEomtLoss.__init__M  sÖ   ø€ õ 	‰Œ×ÒÑÔÐÝ˜$  	Ñ*Ô*Ð*Ø Ô+ˆŒØ&ˆÔð Ô/ˆŒÝ”z $¤/°AÑ"5Ñ6Ô6ˆØœ=ˆ�RÑØ×Ò˜^¨\Ñ:Ô:Ð:ð !Ô1ˆŒØ &Ô 7ˆÔØ'-Ô'EˆÔ$å+ØÔ*ØÔ(ØÔ(Ø”ð	
ñ 
ô 
ˆŒˆˆr/   Úsizesr3   c                 óŒ   — |d         }|dd …         D ]0}t          |¦  «        D ]\  }}t          ||         |¦  «        ||<   ŒŒ1|S )Nr   r   )Ú	enumerateÚmax)rf   r®   ÚmaxesÚsublistÚindexÚitems         r0   Ú_max_by_axiszEomtLoss._max_by_axisp  s`   € Ø�a”ˆØ˜Q˜R˜R”yð 	7ð 	7ˆGÝ(¨Ñ1Ô1ð 7ð 7‘��tÝ" 5¨¤<°Ñ6Ô6��e‘�ð7àˆr/   Útensorsc                 ó"  — |                       d„ |D ¦   «         ¦  «        }t          |¦  «        g|z   }|\  }}}}|d         j        }|d         j        }	t	          j        |||	¬¦  «        }
t	          j        |||ft          j        |	¬¦  «        }t          ||
|¦  «        D ]l\  }}}|d |j	        d         …d |j	        d         …d |j	        d         …f          
                    |¦  «         d|d |j	        d         …d |j	        d         …f<   Œm|
|fS )Nc                 ó6   — g | ]}t          |j        ¦  «        ‘ŒS r.   )r-   rP   )rt   r~   s     r0   rw   z8EomtLoss._pad_images_to_max_in_batch.<locals>.<listcomp>z  s"   € Ð%OÐ%OÐ%O¸V¥d¨6¬<Ñ&8Ô&8Ð%OÐ%OÐ%Or/   r   ©rq   rl   r   r5   F)r¶   Úlenrq   rl   r)   Úzerosr¥   ÚboolÚziprP   Úcopy_)rf   r·   Úmax_sizeÚbatch_shaper„   Ú_ÚheightÚwidthrq   rl   Úpadded_tensorsÚpadding_masksr~   Úpadded_tensorÚpadding_masks                  r0   Ú_pad_images_to_max_in_batchz$EomtLoss._pad_images_to_max_in_batchx  s1  € à×$Ò$Ð%OÐ%OÀwÐ%OÑ%OÔ%OÑPÔPˆå˜7‘|”|�n xÑ/ˆØ'2Ñ$ˆ
�A�v˜uØ˜”
Ô ˆØ˜”Ô"ˆÝœ [¸ÀfÐMÑMÔMˆÝœ
 J°¸Ð#>ÅeÄjÐY_Ð`Ñ`Ô`ˆå36°wÀÐP]Ñ3^Ô3^ð 	Gð 	GÑ/ˆF�M <ØÐ+˜FœL¨œOÐ+Ð->¨v¬|¸A¬Ð->Ð@QÀ&Ä,ÈqÄ/Ð@QÐQÔR×XÒXÐY_Ñ`Ô`Ð`ØAFˆLÐ*˜6œ<¨œ?Ð*Ð,=¨f¬l¸1¬oÐ,=Ð=Ñ>Ð>à˜}Ð,Ð,r/   r   ri   rƒ   c                 óˆ  — |}|j         \  }}}t          j        | j        ¬¦  «        }|                      |¦  «        }	t          j        d„ t          ||¦  «        D ¦   «         ¦  «        }
t          j        ||f| j	        t
          j
        |j        ¬¦  «        }|
||	<   |                     dd¦  «        } |||¦  «        }d|i}|S )a…  Compute the losses related to the labels using cross entropy.

        Args:
            class_queries_logits (`torch.Tensor`):
                A tensor of shape `batch_size, num_queries, num_labels`
            class_labels (`list[torch.Tensor]`):
                List of class labels of shape `(labels)`.
            indices (`tuple[np.array])`:
                The indices computed by the Hungarian matcher.

        Returns:
            `dict[str, Tensor]`: A dict of `torch.Tensor` containing the following key:
            - **loss_cross_entropy** -- The loss computed using cross entropy on the predicted and ground truth labels.
        )Úweightc                 ó*   — g | ]\  }\  }}||         ‘ŒS r.   r.   )rt   ÚtargetrÂ   rv   s       r0   rw   z(EomtLoss.loss_labels.<locals>.<listcomp>Ÿ  s$   € ÐHÐHÐH™>˜6¡6 A qˆV�AŒYÐHÐHÐHr/   )Ú
fill_valuerq   rl   r   r5   Úloss_cross_entropy)rP   r   ÚCrossEntropyLossr¡   Ú$_get_predictions_permutation_indicesr)   Úcatr¾   Úfullr¢   rs   rl   Ú	transpose)rf   r   ri   rƒ   Úpred_logitsr„   Únum_queriesrÂ   rU   ÚidxÚtarget_classes_oÚtarget_classesÚpred_logits_transposedÚloss_ceÚlossess                  r0   Úloss_labelszEomtLoss.loss_labels‰  sß   € ð" +ˆØ%0Ô%6Ñ"ˆ
�K ÝÔ'¨tÔ/@ÐAÑAÔAˆ	Ø×7Ò7¸Ñ@Ô@ˆÝ œ9ØHÐH­S°¸wÑ-GÔ-GÐHÑHÔHñ
ô 
Ðõ œØ˜Ð%°$´/ÍÌÐ]hÔ]oð
ñ 
ô 
ˆð /ˆ�sÑà!,×!6Ò!6°q¸!Ñ!<Ô!<ÐØ�)Ð2°NÑCÔCˆØ&¨Ð0ˆØˆr/   r    rh   r“   c                 óf  ‡ — ‰                       |¦  «        }‰                      |¦  «        }||         }‰                      |¦  «        \  }}	||         }|dd…df         }|dd…df         }t          j        ¦   «         5  ‰                      |ˆ fd„‰ j        ‰ j        ‰ j        ¦  «        }
t          ||
d¬¦  «         
                    d¦  «        }ddd¦  «         n# 1 swxY w Y   t          ||
d¬¦  «         
                    d¦  «        }t          |||¦  «        t          |||¦  «        dœ}~~|S )a¤  Compute the losses related to the masks using sigmoid_cross_entropy_loss and dice loss.

        Args:
            masks_queries_logits (`torch.Tensor`):
                A tensor of shape `(batch_size, num_queries, height, width)`.
            mask_labels (`torch.Tensor`):
                List of mask labels of shape `(labels, height, width)`.
            indices (`tuple[np.array])`:
                The indices computed by the Hungarian matcher.
            num_masks (`int)`:
                The number of masks, used for normalization.

        Returns:
            losses (`dict[str, Tensor]`): A dict of `torch.Tensor` containing two keys:
            - **loss_mask** -- The loss computed using sigmoid cross entropy loss on the predicted and ground truth.
              masks.
            - **loss_dice** -- The loss computed using dice loss on the predicted on the predicted and ground truth,
              masks.
        Nc                 ó.   •— ‰                      | ¦  «        S ©N)Úcalculate_uncertainty)Úlogitsrf   s    €r0   ú<lambda>z%EomtLoss.loss_masks.<locals>.<lambda>Ö  s   ø€ ˜t×9Ò9¸&ÑAÔA€ r/   Frm   r   )Ú	loss_maskÚ	loss_dice)rÑ   Ú _get_targets_permutation_indicesrÉ   r)   r�   Úsample_points_using_uncertaintyra   r¨   r©   r?   r;   rš   r–   )rf   r    rh   rƒ   r“   Úsrc_idxÚtgt_idxÚ
pred_masksÚtarget_masksrÂ   r2   Úpoint_labelsÚpoint_logitsrÜ   s   `             r0   Ú
loss_maskszEomtLoss.loss_masks«  s°  ø€ ð4 ×;Ò;¸GÑDÔDˆØ×7Ò7¸Ñ@Ô@ˆà)¨'Ô2ˆ
ð ×:Ò:¸;ÑGÔG‰ˆ�aØ# GÔ,ˆð      4 Ô(ˆ
Ø# A A A t GÔ,ˆõ Œ]‰_Œ_ð 		ið 		iØ $× DÒ DØØAÐAÐAÐAØ”ØÔ%ØÔ,ñ!ô !Ðõ (¨Ð6GÐW\Ð]Ñ]Ô]×eÒeÐfgÑhÔhˆLð		ið 		ið 		iñ 		iô 		ið 		ið 		ið 		ið 		ið 		ið 		iøøøð 		ið 		ið 		ið 		iõ $ JÐ0AÐQVÐWÑWÔW×_Ò_Ð`aÑbÔbˆõ 4°LÀ,ÐPYÑZÔZÝ" <°¸yÑIÔIð
ð 
ˆð
 ØØˆs   Á?ACÃC Ã#C c                 óœ   — t          j        d„ t          |¦  «        D ¦   «         ¦  «        }t          j        d„ |D ¦   «         ¦  «        }||fS )Nc                 óD   — g | ]\  }\  }}t          j        ||¦  «        ‘ŒS r.   ©r)   Ú	full_like)rt   ru   ÚsrcrÂ   s       r0   rw   zAEomtLoss._get_predictions_permutation_indices.<locals>.<listcomp>ë  s,   € Ð"aÐ"aÐ"a¹{¸qÁ(À3È¥5¤?°3¸Ñ#:Ô#:Ð"aÐ"aÐ"ar/   c                 ó   — g | ]\  }}|‘ŒS r.   r.   )rt   ró   rÂ   s      r0   rw   zAEomtLoss._get_predictions_permutation_indices.<locals>.<listcomp>ì  s   € Ð(EÐ(EÐ(E±°#°q¨Ð(EÐ(EÐ(Er/   ©r)   rÒ   r°   )rf   rƒ   Úbatch_indicesÚpredictions_indicess       r0   rÑ   z-EomtLoss._get_predictions_permutation_indicesé  sT   € åœ	Ð"aÐ"aÍiÐX_ÑN`ÔN`Ð"aÑ"aÔ"aÑbÔbˆÝ#œiÐ(EÐ(E¸WÐ(EÑ(EÔ(EÑFÔFÐØÐ1Ð1Ð1r/   c                 óœ   — t          j        d„ t          |¦  «        D ¦   «         ¦  «        }t          j        d„ |D ¦   «         ¦  «        }||fS )Nc                 óD   — g | ]\  }\  }}t          j        ||¦  «        ‘ŒS r.   rñ   )rt   ru   rÂ   Útgts       r0   rw   z=EomtLoss._get_targets_permutation_indices.<locals>.<listcomp>ñ  s,   € Ð"aÐ"aÐ"a¹{¸qÁ(À1Àc¥5¤?°3¸Ñ#:Ô#:Ð"aÐ"aÐ"ar/   c                 ó   — g | ]\  }}|‘ŒS r.   r.   )rt   rÂ   rú   s      r0   rw   z=EomtLoss._get_targets_permutation_indices.<locals>.<listcomp>ò  s   € Ð#@Ð#@Ð#@©H¨Q° CÐ#@Ð#@Ð#@r/   rõ   )rf   rƒ   rö   Útarget_indicess       r0   ræ   z)EomtLoss._get_targets_permutation_indicesï  sR   € åœ	Ð"aÐ"aÍiÐX_ÑN`ÔN`Ð"aÑ"aÔ"aÑbÔbˆÝœÐ#@Ð#@¸Ð#@Ñ#@Ô#@ÑAÔAˆØ˜nÐ,Ð,r/   râ   c                 ó0   — t          j        |¦  «         }|S )a‚  
        In Eomt paper, uncertainty is estimated as L1 distance between 0.0 and the logit prediction in 'logits'
        for the foreground class in `classes`.

        Args:
            logits (`torch.Tensor`):
            A tensor of shape (R, 1, ...) for class-specific or class-agnostic, where R is the total number of predicted masks in all images and C is:
            the number of foreground classes. The values are logits.

        Returns:
            scores (`torch.Tensor`): A tensor of shape (R, 1, ...) that contains uncertainty scores with the most
            uncertain locations having the highest uncertainty score.
        )r)   Úabs)rf   râ   Úuncertainty_scoress      r0   rá   zEomtLoss.calculate_uncertaintyõ  s   € õ  %œy¨Ñ0Ô0Ð1ÐØ!Ð!r/   ra   r¨   r©   c           	      ó¬  — |j         d         }t          ||z  ¦  «        }t          j        ||d|j        ¬¦  «        }t          ||d¬¦  «        }	 ||	¦  «        }
t          ||z  ¦  «        }||z
  }t          j        |
dd…ddd…f         |d¬¦  «        d         }|t          j        |t          j        |j        ¬	¦  «        z  }||dd…df         z  }| 	                    d
d¦  «        | 	                    d
¦  «        dd…f          	                    ||d¦  «        }|dk    r3t          j
        |t          j        ||d|j        ¬¦  «        gd¬¦  «        }|S )a€  
        This function is meant for sampling points in [0, 1] * [0, 1] coordinate space based on their uncertainty. The
        uncertainty is calculated for each point using the passed `uncertainty function` that takes points logit
        prediction as input.

        Args:
            logits (`float`):
                Logit predictions for P points.
            uncertainty_function:
                A function that takes logit predictions for P points and returns their uncertainties.
            num_points (`int`):
                The number of points P to sample.
            oversample_ratio (`int`):
                Oversampling parameter.
            importance_sample_ratio (`float`):
                Ratio of points that are sampled via importance sampling.

        Returns:
            point_coordinates (`torch.Tensor`):
                Coordinates for P sampled points.
        r   r5   rk   Frm   Nr   )Úkr7   rº   rC   ©r7   )rP   r�   r)   r{   rl   r?   ÚtopkÚarangeÚlongÚviewrÒ   )rf   râ   Úuncertainty_functionra   r¨   r©   Ú	num_boxesÚnum_points_sampledr2   rí   Úpoint_uncertaintiesÚnum_uncertain_pointsÚnum_random_pointsr×   Úshifts                  r0   rç   z(EomtLoss.sample_points_using_uncertainty  s�  € ð< ”L ”Oˆ	Ý  Ð.>Ñ!>Ñ?Ô?Ðõ "œJ yÐ2DÀaÐPVÔP]Ð^Ñ^Ô^Ðå# FÐ,=ÈUÐSÑSÔSˆà2Ð2°<Ñ@Ô@Ðå"Ð#:¸ZÑ#GÑHÔHÐØ&Ð)=Ñ=ÐåŒjÐ,¨Q¨Q¨Q°°1°1°1¨WÔ5Ð9MÐSTÐUÑUÔUÐVWÔXˆØ"¥U¤\°)Å5Ä:ÐV\ÔVcÐ%dÑ%dÔ%dÑdˆØˆu�Q�Q�Q˜�WŒ~ÑˆØ-×2Ò2°2°qÑ9Ô9¸#¿(º(À2¹,¼,ÈÈÈ¸/ÔJ×OÒOÐPYÐ[oÐqrÑsÔsÐà˜qÒ Ð Ý %¤	Ø"¥E¤J¨yÐ:KÈQÐW]ÔWdÐ$eÑ$eÔ$eÐfØð!ñ !ô !Ðð !Ð r/   NÚauxiliary_predictionsc                 óÆ  ‡— |                       ||||¦  «        }|                      ||d         j        ¬¦  «        }i |                      ||||¦  «        ¥|                      |||¦  «        ¥}|�rt          |¦  «        D ]b\  Š}	|	d         }|	d         }|                      ||||¦  «        }
ˆfd„|
                     ¦   «         D ¦   «         }
|                     |
¦  «         Œc|S )a²  
        This performs the loss computation.

        Args:
            masks_queries_logits (`torch.Tensor`):
                A tensor of shape `(batch_size, num_queries, height, width)`.
            class_queries_logits (`torch.Tensor`):
                A tensor of shape `(batch_size, num_queries, num_labels)`.
            mask_labels (`torch.Tensor`):
                List of mask labels of shape `(labels, height, width)`.
            class_labels (`list[torch.Tensor]`):
                List of class labels of shape `(labels)`.
            auxiliary_predictions (`dict[str, torch.Tensor]`, *optional*):
                if `use_auxiliary_loss` was set to `true` in [`EomtConfig`], then it contains the logits from
                the inner layers of the EomtMaskedAttentionDecoder.

        Returns:
            losses (`dict[str, Tensor]`): A dict of `torch.Tensor` containing three keys:
            - **loss_cross_entropy** -- The loss computed using cross entropy on the predicted and ground truth labels.
            - **loss_mask** -- The loss computed using sigmoid cross_entropy loss on the predicted and ground truth
              masks.
            - **loss_dice** -- The loss computed using dice loss on the predicted on the predicted and ground truth
              masks.
            if `use_auxiliary_loss` was set to `true` in [`EomtConfig`], the dictionary contains additional
            losses for each auxiliary predictions.
        r   rk   Nr    r   c                 ó&   •— i | ]\  }}|› d ‰› �|“ŒS )rÂ   r.   )rt   ÚkeyÚvaluer×   s      €r0   ú
<dictcomp>z$EomtLoss.forward.<locals>.<dictcomp>o  s)   ø€ ÐWÐWÐW±z°s¸E ˜^˜^ c˜^˜^¨UÐWÐWÐWr/   )	r­   Úget_num_masksrl   rî   rÝ   r°   r�   ÚitemsÚupdate)rf   r    r   rh   ri   r  rƒ   r“   rÜ   Úaux_outputsÚ	loss_dictr×   s              @r0   r�   zEomtLoss.forward=  s  ø€ ðH —,’,Ð3Ð5IÈ;ÐXdÑeÔeˆà×&Ò& |¸LÈ¼OÔ<RÐ&ÑSÔSˆ	ð%
Ø�oŠoÐ2°KÀÈ)ÑTÔTð%
à×ÒÐ3°\À7ÑKÔKð%
ˆð
 !Ð,Ý$-Ð.CÑ$DÔ$Dð )ð )Ñ ��[Ø'2Ð3IÔ'JÐ$Ø'2Ð3IÔ'JÐ$Ø ŸLšLÐ)=Ð?SÐU`ÐbnÑoÔo�	ØWÐWÐWÐWÀYÇ_Â_ÑEVÔEVÐWÑWÔW�	Ø—’˜iÑ(Ô(Ð(Ð(àˆr/   rl   c                 ó0  — t          d„ |D ¦   «         ¦  «        }t          j        |t          j        |¬¦  «        }d}t	          ¦   «         r2t
          j        i k    r"t          |¦  «        }t          ¦   «         j        }t          j	        ||z  d¬¦  «        }|S )zk
        Computes the average number of target masks across the batch, for normalization purposes.
        c              3   ó4   K  — | ]}t          |¦  «        V — Œd S rà   )r»   )rt   Úclassess     r0   ú	<genexpr>z)EomtLoss.get_num_masks.<locals>.<genexpr>x  s(   è è € ÐAÐA¨�˜G™œÐAÐAÐAÐAÐAÐAr/   rº   r   )Úmin)
rH   r)   rr   rŽ   r   r   Ú_shared_stater   Únum_processesÚclamp)rf   ri   rl   r“   Ú
world_sizes        r0   r  zEomtLoss.get_num_maskst  s‘   € õ ÐAÐA°LÐAÑAÔAÑAÔAˆ	Ý”O IµU´[ÈÐPÑPÔPˆ	Øˆ
Ý"Ñ$Ô$ð 	:ÝÔ)¨RÒ/Ð/Ý" 9Ñ-Ô-�	Ý)™^œ^Ô9�
å”K 	¨JÑ 6¸AÐ>Ñ>Ô>ˆ	ØÐr/   rà   )r%   r&   r'   r   ÚdictÚstrrŽ   rd   r-   r�   r¶   r   r,   rÉ   ÚnpÚarrayrÝ   r)   rî   rÑ   ræ   rá   rç   r�   rl   r  r‘   r’   s   @r0   rœ   rœ   L  s¦  ø€ € € € € ð!
˜zð !
¸¸SÀ%¸ZÔ8Hð !
ð !
ð !
ð !
ð !
ð !
ðF $ t¨C¤y¤/ð °d¸3´ið ð ð ð ð-°4¸´<ð -ÀEÈ&ÐRXÈ.ÔDYð -ð -ð -ð -ð" Ø$*ð Ø:>¸v¼,ð ØQVÐWYÔW_ÔQ`ð à	ˆc�6ˆkÔ	ð ð  ð  ð  ðD<à#œlð<ð ˜%œ,Ô'ð<ð �r”x”ð	<ð
 ð<ð 
ˆc�5”<ÐÔ	 ð<ð <ð <ð <ð|2ð 2ð 2ð-ð -ð -ð"¨E¬Lð "¸U¼\ð "ð "ð "ð "ð"5!à”ð5!ð ð	5!ð
 ð5!ð "'ð5!ð 
Œð5!ð 5!ð 5!ð 5!ðz AEð5ð 5à#œlð5ð $œlð5ð ˜%œ,Ô'ð	5ð
 ˜5œ<Ô(ð5ð  $ C¨¬Ð$5Ô6¸Ñ=ð5ð 
ˆc�5”<ÐÔ	 ð5ð 5ð 5ð 5ðn¨%¬,ð ÀÄð ÐQVÔQ]ð ð ð ð ð ð ð ð r/   rœ   c                   óF   ‡ — e Zd ZdZˆ fd„Zdej        dej        fd„Zˆ xZS )ÚEomtPatchEmbeddingszì
    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.
    c                 óÌ  •— t          ¦   «                              ¦   «          |j        |j        }}|j        |j        }}t          |t          j        j	        ¦  «        r|n||f}t          |t          j        j	        ¦  «        r|n||f}|d         |d         z  |d         |d         z  z  }|| _        || _        || _        || _
        t          j        ||||¬¦  «        | _        d S )Nr   r   ©Úkernel_sizeÚstride)rc   rd   Ú
image_sizeÚ
patch_sizeÚnum_channelsÚhidden_sizeÚ
isinstanceÚcollectionsÚabcÚIterableÚnum_patchesr   ÚConv2dÚ
projection)rf   r�   r,  r-  r.  r/  r4  rg   s          €r0   rd   zEomtPatchEmbeddings.__init__‹  sá   ø€ Ý‰Œ×ÒÑÔÐØ!'Ô!2°FÔ4E�Jˆ
Ø$*Ô$7¸Ô9K�kˆå#-¨j½+¼/Ô:RÑ#SÔ#SÐq�Z�ZÐZdÐfpÐYqˆ
Ý#-¨j½+¼/Ô:RÑ#SÔ#SÐq�Z�ZÐZdÐfpÐYqˆ
Ø! !”}¨
°1¬Ñ5¸*ÀQ¼-È:ÐVWÌ=Ñ:XÑYˆØ$ˆŒØ$ˆŒØ(ˆÔØ&ˆÔåœ) L°+È:Ð^hÐiÑiÔiˆŒˆˆr/   Úpixel_valuesr3   c                 óä   — |j         d         }|| j        k    rt          d| j        › d|› d�¦  «        ‚|                      |¦  «                             d¦  «                             dd¦  «        }|S )Nr   zoMake sure that the channel dimension of the pixel values match with the one set in the configuration. Expected z	 but got ú.r5   )rP   r.  re   r6  rE   rÔ   )rf   r7  r.  Ú
embeddingss       r0   r�   zEomtPatchEmbeddings.forwardš  s“   € Ø#Ô)¨!Ô,ˆØ˜4Ô,Ò,Ð,ÝðIØ!Ô.ðIð IØ9EðIð Ið Iñô ð ð —_’_ \Ñ2Ô2×:Ò:¸1Ñ=Ô=×GÒGÈÈ1ÑMÔMˆ
ØÐr/   )	r%   r&   r'   r(   rd   r)   r   r�   r‘   r’   s   @r0   r'  r'  „  sm   ø€ € € € € ðð ðjð jð jð jð jð E¤Lð °U´\ð ð ð ð ð ð ð ð r/   r'  c                   óP   ‡ — e Zd ZdZdeddfˆ fd„Zdej        dej        fd„Zˆ xZ	S )ÚEomtEmbeddingszM
    Construct the CLS token, mask token, position and patch embeddings.
    r�   r3   Nc                 ó’  •— t          ¦   «                              ¦   «          || _        |j        | _        t	          j        t          j        dd|j        ¦  «        ¦  «        | _	        t	          j        t          j
        d|j        |j        ¦  «        ¦  «        | _        t          |¦  «        | _        | j        j        }t	          j        |j        ¦  «        | _        d|j        z   | _        t	          j        ||j        ¦  «        | _        |                      dt          j        |¦  «                             d¦  «        d¬¦  «         d S )Nr   Úposition_ids©r   rC   F)Ú
persistent)rc   rd   r�   r-  r   Ú	Parameterr)   Úrandnr/  Ú	cls_tokenr¼   Únum_register_tokensÚregister_tokensr'  Úpatch_embeddingsr4  ÚDropoutÚhidden_dropout_probÚdropoutÚnum_prefix_tokensÚ	EmbeddingÚposition_embeddingsr¦   r  Úexpand)rf   r�   r4  rg   s      €r0   rd   zEomtEmbeddings.__init__ª  s  ø€ Ý‰Œ×ÒÑÔÐàˆŒØ Ô+ˆŒåœ¥e¤k°!°Q¸Ô8JÑ&KÔ&KÑLÔLˆŒÝ!œ|­E¬K¸¸6Ô;UÐW]ÔWiÑ,jÔ,jÑkÔkˆÔå 3°FÑ ;Ô ;ˆÔØÔ+Ô7ˆÝ”z &Ô"<Ñ=Ô=ˆŒØ!" VÔ%?Ñ!?ˆÔÝ#%¤<°¸VÔ=OÑ#PÔ#PˆÔ Ø×Ò˜^­U¬\¸+Ñ-FÔ-F×-MÒ-MÈgÑ-VÔ-VÐchÐÑiÔiÐiÐiÐir/   r7  c                 ó¢  — |j         \  }}}}| j        j        j        j        }|                      |                     |¬¦  «        ¦  «        }| j                             |dd¦  «        }| j                             |dd¦  «        }||  	                    | j
        ¦  «        z   }t          j        |||gd¬¦  «        }|                      |¦  «        }|S )Nrp   rC   r   r  )rP   rF  r6  rË   rq   rz   rC  rM  rE  rL  r>  r)   rÒ   rI  )rf   r7  r„   rÂ   Útarget_dtyper:  Ú
cls_tokensrE  s           r0   r�   zEomtEmbeddings.forwardº  sÅ   € Ø*Ô0Ñˆ
�A�q˜!ØÔ,Ô7Ô>ÔDˆØ×*Ò*¨<¯?ª?À¨?Ñ+NÔ+NÑOÔOˆ
à”^×*Ò*¨:°r¸2Ñ>Ô>ˆ
ØÔ.×5Ò5°jÀ"ÀbÑIÔIˆà $×":Ò":¸4Ô;LÑ"MÔ"MÑMˆ
Ý”Y 
¨O¸ZÐHÈaÐPÑPÔPˆ
à—\’\ *Ñ-Ô-ˆ
àÐr/   ©
r%   r&   r'   r(   r   rd   r)   r   r�   r‘   r’   s   @r0   r<  r<  ¥  sƒ   ø€ € € € € ðð ðj˜zð j¨dð jð jð jð jð jð jð  E¤Lð °U´\ð ð ð ð ð ð ð ð r/   r<  ç        ÚmoduleÚqueryr  r  Úattention_maskÚscalingrI  c                 óÀ  — t          j        ||                     dd¦  «        ¦  «        |z  }|�||z   }t          j                             |dt           j        ¬¦  «                             |j        ¦  «        }t          j         	                    ||| j
        ¬¦  «        }t          j        ||¦  «        }	|	                     dd¦  «                             ¦   «         }	|	|fS )NrC   éþÿÿÿ)r7   rq   )ÚpÚtrainingr   r5   )r)   rF   rÔ   r   r9   ry   Úfloat32rz   rq   rI  rZ  Ú
contiguous)
rS  rT  r  r  rU  rV  rI  r=   Úattn_weightsÚattn_outputs
             r0   Úeager_attention_forwardr_  Ê  sÃ   € õ ”<  s§}¢}°R¸Ñ'<Ô'<Ñ=Ô=ÀÑG€LØÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2ÅUÄ]Ð(ÑSÔS×VÒVÐW\ÔWbÑcÔc€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€Lå”,˜|¨UÑ3Ô3€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r/   c            
       ó~   ‡ — e Zd ZdZˆ fd„Z	 ddej        dej        dz  deej        ej        dz  f         fd„Zˆ xZ	S )	ÚEomtAttentionz=Multi-headed attention from 'Attention Is All You Need' paperc                 ó‚  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        | j        | j        z  | _        | j        | j        z  | j        k    r t          d| j        › d| j        › d�¦  «        ‚| j        dz  | _	        |j
        | _        d| _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        d S )Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: z).g      à¿F)rc   rd   r�   r/  Ú	embed_dimÚnum_attention_headsÚ	num_headsÚhead_dimre   ÚscaleÚattention_dropoutrI  Ú	is_causalr   ÚLinearÚk_projÚv_projÚq_projÚout_proj©rf   r�   rg   s     €r0   rd   zEomtAttention.__init__ä  s  ø€ Ý‰Œ×ÒÑÔÐØˆŒØÔ+ˆŒØÔ3ˆŒØœ¨$¬.Ñ8ˆŒØŒ=˜4œ>Ñ)¨T¬^Ò;Ð;Ýð'ÈdÌnð 'ð 'Ø”Nð'ð 'ð 'ñô ð ð ”] DÑ(ˆŒ
ØÔ/ˆŒØˆŒå”i ¤°´Ñ?Ô?ˆŒÝ”i ¤°´Ñ?Ô?ˆŒÝ”i ¤°´Ñ?Ô?ˆŒÝœ	 $¤.°$´.ÑAÔAˆŒˆˆr/   Nr"   rU  r3   c           
      ó¼  — |j         dd…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }t          j        | j	        j
        t          ¦  «        }	 |	| ||||| j        | j        | j        sdn| j        ¬¦  «        \  }
} |
j        g |¢d‘R Ž                      ¦   «         }
|                      |
¦  «        }
|
|fS )z#Input shape: Batch x Time x ChannelNrC   r   r5   rR  )ri  rV  rI  )rP   rf  rm  r  rÔ   rk  rl  r   Úget_interfacer�   Ú_attn_implementationr_  ri  rg  rZ  rI  Úreshaper\  rn  )rf   r"   rU  r=   Úinput_shapeÚhidden_shapeÚqueriesÚkeysÚvaluesÚattention_interfacer^  r]  s               r0   r�   zEomtAttention.forwardø  sg  € ð $Ô)¨#¨2¨#Ô.ˆà8˜Ð8 bÐ8¨$¬-Ð8Ð8ˆØ—+’+˜mÑ,Ô,×1Ò1°,Ñ?Ô?×IÒIÈ!ÈQÑOÔOˆØ�{Š{˜=Ñ)Ô)×.Ò.¨|Ñ<Ô<×FÒFÀqÈ!ÑLÔLˆØ—’˜]Ñ+Ô+×0Ò0°Ñ>Ô>×HÒHÈÈAÑNÔNˆå(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð %8Ð$7ØØØØØØ”nØ”JØ#œ}Ð>�C�C°$´,ð	%
ñ 	%
ô 	%
Ñ!ˆ�\ð *�kÔ)Ð;¨;Ð;¸Ð;Ð;Ð;×FÒFÑHÔHˆØ—m’m KÑ0Ô0ˆà˜LÐ(Ð(r/   rà   )
r%   r&   r'   r(   rd   r)   r   r,   r�   r‘   r’   s   @r0   ra  ra  á  s™   ø€ € € € € ØGÐGðBð Bð Bð Bð Bð. /3ð!)ð !)à”|ð!)ð œ tÑ+ð!)ð
 
ˆuŒ|˜Uœ\¨DÑ0Ð0Ô	1ð!)ð !)ð !)ð !)ð !)ð !)ð !)ð !)r/   ra  c                   óD   ‡ — e Zd Zdˆ fd„Zdej        dej        fd„Zˆ xZS )ÚEomtLayerScaler3   Nc                 ó¸   •— t          ¦   «                              ¦   «          t          j        |j        t          j        |j        ¦  «        z  ¦  «        | _        d S rà   )	rc   rd   r   rA  Úlayerscale_valuer)   r¥   r/  Úlambda1ro  s     €r0   rd   zEomtLayerScale.__init__  sC   ø€ Ý‰Œ×ÒÑÔÐÝ”| FÔ$;½e¼jÈÔI[Ñ>\Ô>\Ñ$\Ñ]Ô]ˆŒˆˆr/   Úhidden_statec                 ó   — || j         z  S rà   )r~  ©rf   r  s     r0   r�   zEomtLayerScale.forward!  s   € Ø˜dœlÑ*Ð*r/   ©r3   N©r%   r&   r'   rd   r)   r   r�   r‘   r’   s   @r0   r{  r{    si   ø€ € € € € ð^ð ^ð ^ð ^ð ^ð ^ð+ E¤Lð +°U´\ð +ð +ð +ð +ð +ð +ð +ð +r/   r{  c                   óD   ‡ — e Zd Zdˆ fd„Zdej        dej        fd„Zˆ xZS )ÚEomtMLPr3   Nc                 ó~  •— t          ¦   «                              ¦   «          |j        x}}t          |j        |j        z  ¦  «        }t          j        ||d¬¦  «        | _        t          |j	        t          ¦  «        rt          |j	                 | _        n|j	        | _        t          j        ||d¬¦  «        | _        d S )NT©Úbias)rc   rd   r/  r�   Ú	mlp_ratior   rj  Úfc1r0  Ú
hidden_actr#  r	   Ú
activationÚfc2©rf   r�   Úin_featuresÚout_featuresÚhidden_featuresrg   s        €r0   rd   zEomtMLP.__init__&  s¢   ø€ Ý‰Œ×ÒÑÔÐØ%+Ô%7Ð7ˆ�lÝ˜fÔ0°6Ô3CÑCÑDÔDˆÝ”9˜[¨/ÀÐEÑEÔEˆŒÝ�fÔ'­Ñ-Ô-ð 	0Ý$ VÔ%6Ô7ˆDŒOˆOà$Ô/ˆDŒOÝ”9˜_¨lÀÐFÑFÔFˆŒˆˆr/   r  c                 ó„   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|S rà   )rŠ  rŒ  r�  r�  s     r0   r�   zEomtMLP.forward1  s;   € Ø—x’x Ñ-Ô-ˆØ—’ |Ñ4Ô4ˆØ—x’x Ñ-Ô-ˆØÐr/   r‚  rƒ  r’   s   @r0   r…  r…  %  si   ø€ € € € € ð	Gð 	Gð 	Gð 	Gð 	Gð 	Gð E¤Lð °U´\ð ð ð ð ð ð ð ð r/   r…  c                   óD   ‡ — e Zd Zdˆ fd„Zdej        dej        fd„Zˆ xZS )ÚEomtSwiGLUFFNr3   Nc                 óD  •— t          ¦   «                              ¦   «          |j        x}}t          |j        |j        z  ¦  «        }t          |dz  dz  ¦  «        dz   dz  dz  }t          j        |d|z  d¬¦  «        | _        t          j        ||d¬¦  «        | _        d S )Nr5   r   é   é   Tr‡  )	rc   rd   r/  r�   r‰  r   rj  Ú
weights_inÚweights_outrŽ  s        €r0   rd   zEomtSwiGLUFFN.__init__9  s    ø€ Ý‰Œ×ÒÑÔÐØ%+Ô%7Ð7ˆ�lÝ˜fÔ0°6Ô3CÑCÑDÔDˆÝ˜°Ñ2°QÑ6Ñ7Ô7¸!Ñ;ÀÑAÀAÑEˆåœ) K°°_Ñ1DÈ4ÐPÑPÔPˆŒÝœ9 _°lÈÐNÑNÔNˆÔÐÐr/   r  c                 óÎ   — |                       |¦  «        }|                     dd¬¦  «        \  }}t          j                             |¦  «        |z  }|                      |¦  «        S )Nr5   rC   r  )r˜  Úchunkr   r9   Úsilur™  )rf   r  Úx1Úx2Úhiddens        r0   r�   zEomtSwiGLUFFN.forwardB  s]   € Ø—’ |Ñ4Ô4ˆØ×#Ò# A¨2Ð#Ñ.Ô.‰ˆˆBÝ”×#Ò# BÑ'Ô'¨"Ñ,ˆØ×Ò Ñ'Ô'Ð'r/   r‚  rƒ  r’   s   @r0   r”  r”  8  si   ø€ € € € € ðOð Oð Oð Oð Oð Oð( E¤Lð (°U´\ð (ð (ð (ð (ð (ð (ð (ð (r/   r”  c                   ó^   ‡ — e Zd ZdZd
deddfˆ fd„Zdej        dej        fd„Zde	fd	„Z
ˆ xZS )ÚEomtDropPathzÏStochastic depth (DropPath) per sample, for residual blocks.

    Identity when ``drop_prob`` is 0 or outside training. See `Deep Networks with Stochastic Depth
    <https://arxiv.org/abs/1603.09382>`_.
    rR  Ú	drop_probr3   Nc                 óV   •— t          ¦   «                              ¦   «          || _        d S rà   )rc   rd   r¢  )rf   r¢  rg   s     €r0   rd   zEomtDropPath.__init__P  s$   ø€ Ý‰Œ×ÒÑÔÐØ"ˆŒˆˆr/   r"   c                 ó  — | j         dk    s| j        s|S d| j         z
  }|j        d         fd|j        dz
  z  z   }t	          j        ||j        |j        ¬¦  «        }t	          j        ||z   ¦  «        }| 	                    |¦  «        |z  S )NrR  r   r   )r   rº   )
r¢  rZ  rP   Úndimr)   r{   rq   rl   ÚfloorÚdiv)rf   r"   Ú	keep_probrP   Úrandom_tensors        r0   r�   zEomtDropPath.forwardT  s“   € ØŒ>˜SÒ Ð ¨¬Ð Ø Ð Ø˜œÑ&ˆ	ØÔ$ QÔ'Ð)¨D°MÔ4FÈÑ4JÑ,KÑKˆÝœ
 5°Ô0CÈMÔL`ÐaÑaÔaˆÝœ M°IÑ$=Ñ>Ô>ˆØ× Ò  Ñ+Ô+¨mÑ;Ð;r/   c                 ó   — d| j         › �S )Nzp=)r¢  ©rf   s    r0   Ú
extra_reprzEomtDropPath.extra_repr]  s   € Ø$�D”NÐ$Ð$Ð$r/   ©rR  )r%   r&   r'   r(   rŽ   rd   r)   r   r�   r#  r¬  r‘   r’   s   @r0   r¡  r¡  I  s›   ø€ € € € € ðð ð#ð # %ð #°$ð #ð #ð #ð #ð #ð #ð< U¤\ð <°e´lð <ð <ð <ð <ð%˜Cð %ð %ð %ð %ð %ð %ð %ð %r/   r¡  c                   óh   ‡ — e Zd ZdZdeddfˆ fd„Z	 d	dej        dej        dz  dej        fd„Zˆ xZ	S )
Ú	EomtLayerzCThis corresponds to the Block class in the original implementation.r�   r3   Nc                 ó"  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¬¦  «        | _        t          |¦  «        | _        t          |¦  «        | _
        |j        dk    rt          |j        ¦  «        nt          j        ¦   «         | _        t          j        |j        |j        ¬¦  «        | _        |j        rt#          |¦  «        | _        nt'          |¦  «        | _        t          |¦  «        | _        d S )N©ÚepsrR  )rc   rd   r   Ú	LayerNormr/  Úlayer_norm_epsÚnorm1ra  Ú	attentionr{  Úlayer_scale1Údrop_path_rater¡  ÚIdentityÚ	drop_pathÚnorm2Úuse_swiglu_ffnr”  Úmlpr…  Úlayer_scale2ro  s     €r0   rd   zEomtLayer.__init__d  sÞ   ø€ Ý‰Œ×ÒÑÔÐå”\ &Ô"4¸&Ô:OÐPÑPÔPˆŒ
Ý& vÑ.Ô.ˆŒÝ*¨6Ñ2Ô2ˆÔØ@FÔ@UÐX[Ò@[Ð@[� fÔ&;Ñ<Ô<Ð<ÕacÔalÑanÔanˆŒå”\ &Ô"4¸&Ô:OÐPÑPÔPˆŒ
àÔ ð 	'Ý$ VÑ,Ô,ˆDŒHˆHå˜v‘”ˆDŒHÝ*¨6Ñ2Ô2ˆÔÐÐr/   r"   rU  c                 ój  — |                       |¦  «        }|                      ||¦  «        \  }}|                      |¦  «        }|                      |¦  «        |z   }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        |z   }|S rà   )rµ  r¶  r·  rº  r»  r½  r¾  )rf   r"   rU  Úhidden_states_normÚself_attention_outputrÂ   Úlayer_outputs          r0   r�   zEomtLayer.forwardt  s³   € ð
 "ŸZšZ¨Ñ6Ô6ÐØ#'§>¢>Ð2DÀnÑ#UÔ#UÑ Ð˜qØ $× 1Ò 1Ð2GÑ HÔ HÐð ŸšÐ'<Ñ=Ô=ÀÑMˆð —z’z -Ñ0Ô0ˆØ—x’x Ñ-Ô-ˆØ×(Ò(¨Ñ6Ô6ˆð —~’~ lÑ3Ô3°mÑCˆàÐr/   rà   rQ  r’   s   @r0   r¯  r¯  a  s–   ø€ € € € € ØMÐMð3˜zð 3¨dð 3ð 3ð 3ð 3ð 3ð 3ð& /3ðð à”|ðð œ tÑ+ðð 
Œð	ð ð ð ð ð ð ð r/   r¯  c                   óD   ‡ — e Zd Zdˆ fd„	Zdej        dej        fd„Zˆ xZS )ÚEomtLayerNorm2dç�íµ ÷Æ°>Tc                 óP   •— t          ¦   «                              |||¬¦  «         d S )N)r²  Úelementwise_affine)rc   rd   )rf   r.  r²  Úaffinerg   s       €r0   rd   zEomtLayerNorm2d.__init__Œ  s(   ø€ Ý‰Œ×Ò˜¨3À6ÐÑJÔJÐJÐJÐJr/   r  r3   c                 ó¾   — |                      dddd¦  «        }t          j        || j        | j        | j        | j        ¦  «        }|                      dddd¦  «        }|S )Nr   r5   r   r   )ÚpermuteÚFÚ
layer_normÚnormalized_shaperË   rˆ  r²  r�  s     r0   r�   zEomtLayerNorm2d.forward�  s^   € Ø#×+Ò+¨A¨q°!°QÑ7Ô7ˆÝ”| L°$Ô2GÈÌÐVZÔV_ÐaeÔaiÑjÔjˆØ#×+Ò+¨A¨q°!°QÑ7Ô7ˆØÐr/   )rÅ  Trƒ  r’   s   @r0   rÄ  rÄ  ‹  si   ø€ € € € € ðKð Kð Kð Kð Kð Kð E¤Lð °U´\ð ð ð ð ð ð ð ð r/   rÄ  c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚEomtScaleLayerr�   c                 ó$  •— t          ¦   «                              ¦   «          |j        }t          j        ||dd¬¦  «        | _        t          |j                 | _        t          j	        ||dd|d¬¦  «        | _
        t          |¦  «        | _        d S )Nr5   r)  r   r   F)r*  ÚpaddingÚgroupsrˆ  )rc   rd   r/  r   ÚConvTranspose2dÚconv1r	   r‹  rŒ  r5  Úconv2rÄ  Úlayernorm2d©rf   r�   r/  rg   s      €r0   rd   zEomtScaleLayer.__init__—  sŽ   ø€ Ý‰Œ×ÒÑÔÐØÔ(ˆÝÔ'¨°[ÈaÐXYÐZÑZÔZˆŒ
Ý  Ô!2Ô3ˆŒÝ”YØØØØØØð
ñ 
ô 
ˆŒ
õ +¨;Ñ7Ô7ˆÔÐÐr/   r"   r3   c                 ó®   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|S rà   )rÔ  rŒ  rÕ  rÖ  ©rf   r"   s     r0   r�   zEomtScaleLayer.forward§  sN   € ØŸ
š
 =Ñ1Ô1ˆØŸš¨Ñ6Ô6ˆØŸ
š
 =Ñ1Ô1ˆØ×(Ò(¨Ñ7Ô7ˆØÐr/   ©	r%   r&   r'   r   rd   r)   r   r�   r‘   r’   s   @r0   rÏ  rÏ  –  sj   ø€ € € € € ð8˜zð 8ð 8ð 8ð 8ð 8ð 8ð  U¤\ð °e´lð ð ð ð ð ð ð ð r/   rÏ  c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚEomtScaleBlockr�   c                 óÐ   •‡— t          ¦   «                              ¦   «          ‰j        | _        t	          j        ˆfd„t          | j        ¦  «        D ¦   «         ¦  «        | _        d S )Nc                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r.   )rÏ  ©rt   rÂ   r�   s     €r0   rw   z+EomtScaleBlock.__init__.<locals>.<listcomp>³  s!   ø€ Ð#[Ð#[Ð#[¸q¥N°6Ñ$:Ô$:Ð#[Ð#[Ð#[r/   )rc   rd   Únum_upscale_blocksÚ
num_blocksr   Ú
ModuleListrx   Úblockro  s    `€r0   rd   zEomtScaleBlock.__init__°  sX   øø€ Ý‰Œ×ÒÑÔÐØ Ô3ˆŒÝ”]Ð#[Ð#[Ð#[Ð#[ÅEÈ$Ì/ÑDZÔDZÐ#[Ñ#[Ô#[Ñ\Ô\ˆŒ
ˆ
ˆ
r/   r"   r3   c                 ó0   — | j         D ]} ||¦  «        }Œ|S rà   )rã  )rf   r"   rã  s      r0   r�   zEomtScaleBlock.forwardµ  s*   € Ø”Zð 	1ð 	1ˆEØ!˜E -Ñ0Ô0ˆMˆMØÐr/   rÚ  r’   s   @r0   rÜ  rÜ  ¯  sq   ø€ € € € € ð]˜zð ]ð ]ð ]ð ]ð ]ð ]ð
 U¤\ð °e´lð ð ð ð ð ð ð ð r/   rÜ  c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚEomtMaskHeadr�   c                 ó   •— t          ¦   «                              ¦   «          |j        }t          j        ||¦  «        | _        t          j        ||¦  «        | _        t          j        ||¦  «        | _        t          |j	                 | _
        d S rà   )rc   rd   r/  r   rj  rŠ  r�  Úfc3r	   r‹  rŒ  r×  s      €r0   rd   zEomtMaskHead.__init__¼  sm   ø€ Ý‰Œ×ÒÑÔÐàÔ(ˆÝ”9˜[¨+Ñ6Ô6ˆŒÝ”9˜[¨+Ñ6Ô6ˆŒÝ”9˜[¨+Ñ6Ô6ˆŒÝ  Ô!2Ô3ˆŒˆˆr/   r"   r3   c                 óÐ   — |                       |                      |¦  «        ¦  «        }|                       |                      |¦  «        ¦  «        }|                      |¦  «        }|S rà   )rŒ  rŠ  r�  rè  rÙ  s     r0   r�   zEomtMaskHead.forwardÅ  sS   € ØŸš¨¯ª°Ñ(?Ô(?Ñ@Ô@ˆØŸš¨¯ª°Ñ(?Ô(?Ñ@Ô@ˆØŸš Ñ/Ô/ˆØÐr/   rÚ  r’   s   @r0   ræ  ræ  »  sj   ø€ € € € € ð4˜zð 4ð 4ð 4ð 4ð 4ð 4ð U¤\ð °e´lð ð ð ð ð ð ð ð r/   ræ  c                   ó�   ‡ — e Zd ZU dZeed<   dZdZdZdZ	dgZ
dZeed	œZ ej        ¦   «         d
ej        ddfˆ fd„¦   «         Zˆ xZS )ÚEomtPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    r�   Úeomtr7  )ÚimageFr¯  T)r"   r#   rS  r3   Nc                 óÚ  •— t          ¦   «                              |¦  «         | j        j        }t	          |t
          j        t
          j        t
          j        f¦  «        rŸt          j
        |j        t          j        d¦  «        ¬¦  «         |j        �it          j        j	                             |j        ¦  «        \  }}|dk    rdt          j        |¦  «        z  nd}t          j        |j        | |¦  «         d S d S t	          |t
          j        ¦  «        rct          j        |j        dd¬¦  «         |j        �<t+          |j        dd¦  «        s(t          j        |j        |j                 ¦  «         d S d S d S t	          |t.          ¦  «        r8t1          |d	¦  «        r&t          j        |j        | j        j        ¦  «         d S d S t	          |t8          ¦  «        r†t          j        |j        d|¬¦  «         t          j        |j        ¦  «         t          j         |j!        t          j"        |j!        j#        d
         ¦  «         $                    d¦  «        ¦  «         d S t	          |tJ          ¦  «        rBt          j&        |j'        dz   ¦  «        }|j(        |d
<   t          j         |j)        |¦  «         d S t	          |tT          ¦  «        rt          j+        |j,        ¦  «         d S d S )Né   )Úar   r   rR  )r˜   ÚstdÚ_is_hf_initializedFr~  rC   r?  )-rc   Ú_init_weightsr�   Úinitializer_ranger0  r   rj  r5  rÓ  ÚinitÚkaiming_uniform_rË   ÚmathÚsqrtrˆ  r)   Ú_calculate_fan_in_and_fan_outÚuniform_rK  Únormal_Úpadding_idxÚgetattrÚzeros_r{  ÚhasattrÚ	constant_r~  r}  r<  Útrunc_normal_rC  rE  r¿   r>  r  rP   rM  rœ   r¥   r¢   r¤   r¡   ÚEomtForUniversalSegmentationÚones_Úattn_mask_probs)rf   rS  rñ  Úfan_inrÂ   Úboundr¡   rg   s          €r0   ró  z!EomtPreTrainedModel._init_weightsß  s¢  ø€ å‰Œ×Ò˜fÑ%Ô%Ð%ØŒkÔ+ˆÝ�f�rœy­"¬)µRÔ5GÐHÑIÔIð 	/ÝÔ! &¤-µ4´9¸Q±<´<Ð@Ñ@Ô@Ð@ØŒ{Ð&Ý!œHœM×GÒGÈÌÑVÔV‘	�˜Ø17¸!²°˜�DœI fÑ-Ô-Ñ-Ð-À�Ý”˜fœk¨E¨6°5Ñ9Ô9Ð9Ð9Ð9ð 'Ð&õ ˜¥¤Ñ-Ô-ð 	/ÝŒL˜œ¨S°aÐ8Ñ8Ô8Ð8àÔ!Ð-µg¸f¼mÐMaÐchÑ6iÔ6iÐ-Ý”˜FœM¨&Ô*<Ô=Ñ>Ô>Ð>Ð>Ð>ð .Ð-Ð-Ð-å˜¥Ñ/Ô/ð 	/Ý�v˜yÑ)Ô)ð MÝ”˜vœ~¨t¬{Ô/KÑLÔLÐLÐLÐLðMð Må˜¥Ñ/Ô/ð 		/ÝÔ˜vÔ/°c¸sÐCÑCÔCÐCÝŒK˜Ô.Ñ/Ô/Ð/ÝŒJ�vÔ*­E¬L¸Ô9LÔ9RÐSUÔ9VÑ,WÔ,W×,^Ò,^Ð_fÑ,gÔ,gÑhÔhÐhÐhÐhÝ˜¥Ñ)Ô)ð 	/Ý œ: fÔ&7¸!Ñ&;Ñ<Ô<ˆLØ%œˆL˜ÑÝŒJ�vÔ*¨LÑ9Ô9Ð9Ð9Ð9Ý˜Õ <Ñ=Ô=ð 	/ÝŒJ�vÔ-Ñ.Ô.Ð.Ð.Ð.ð	/ð 	/r/   )r%   r&   r'   r(   r   r+   Úbase_model_prefixÚmain_input_nameÚinput_modalitiesÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_supports_sdpar¯  ra  Ú_can_record_outputsr)   r�   r   ÚModuleró  r‘   r’   s   @r0   rë  rë  Ì  s²   ø€ € € € € € ðð ð
 ÐÐÑØÐØ$€OØ!ÐØ&+Ð#Ø$˜ÐØ€Nà"Ø#ðð Ðð
 €U„]�_„_ð/ B¤Ið /°$ð /ð /ð /ð /ð /ñ „_ð/ð /ð /ð /ð /r/   rë  zV
    The EoMT Model with head on top for instance/semantic/panoptic segmentation.
    c                   óT  ‡ — e Zd ZdZdefˆ fd„Zdededededeeef         d	eeef         fd
„Z	deeef         d	efd„Z
eee	 	 	 ddedee         dz  dee         dz  dee         dz  dee         d	efd„¦   «         ¦   «         ¦   «         Zd„ Zdej        fd„Zed„ ¦   «         Zˆ xZS )r  r7  r�   c                 ój  •‡— t          ¦   «                              ‰¦  «         ‰| _        ‰j        | _        t	          ‰¦  «        | _        t          j        ‰j        ‰j	        ¬¦  «        | _
        t          j        ‰j        ‰j        ¦  «        | _        t          j        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        t#          ‰¦  «        | _        t'          ‰¦  «        | _        t          j        ‰j        ‰j        dz   ¦  «        | _        ‰j        ‰j        z  ‰j        ‰j        z  f| _        ‰j        ‰j        ‰j        dœ| _        t?          ‰| j        ¬¦  «        | _         |  !                    dtE          j#        ‰j$        ¦  «        ¦  «         |  %                    ¦   «          d S )Nr±  c                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r.   )r¯  rß  s     €r0   rw   z9EomtForUniversalSegmentation.__init__.<locals>.<listcomp>  s!   ø€ Ð$`Ð$`Ð$`¸1¥Y¨vÑ%6Ô%6Ð$`Ð$`Ð$`r/   r   )rÏ   rä   rå   )r�   rž   r  )&rc   rd   r�   Únum_hidden_layersr<  r:  r   r³  r/  r´  Ú	layernormrK  rÖ   rT  râ  rx   ÚlayersrÜ  Úupscale_blockræ  Ú	mask_headrj  r¢   Úclass_predictorr,  r-  Ú	grid_sizerª   r¬   r«   rž   rœ   rU   r¦   r)   r¥   rá  Ú	post_initro  s    `€r0   rd   z%EomtForUniversalSegmentation.__init__  sy  øø€ Ý‰Œ×Ò˜Ñ Ô Ð ØˆŒØ!'Ô!9ˆÔÝ(¨Ñ0Ô0ˆŒÝœ fÔ&8¸fÔ>SÐTÑTÔTˆŒå”\ &Ô"4°fÔ6HÑIÔIˆŒ
Ý”mÐ$`Ð$`Ð$`Ð$`ÅÀfÔF^Ñ@_Ô@_Ð$`Ñ$`Ô$`ÑaÔaˆŒå+¨FÑ3Ô3ˆÔÝ% fÑ-Ô-ˆŒå!œy¨Ô);¸VÔ=NÐQRÑ=RÑSÔSˆÔà Ô+¨vÔ/@Ñ@À&ÔBSÐW]ÔWhÑBhÐiˆŒà"(Ô"5ØÔ+ØÔ+ð.
ð .
ˆÔõ "¨¸TÔ=MÐNÑNÔNˆŒà×ÒÐ.µ´
¸6Ô;LÑ0MÔ0MÑNÔNÐNà�ŠÑÔÐÐÐr/   r    r   rh   ri   r  r3   c                 ó¾   — |                       |||||¬¦  «        }| j                             ¦   «         D ](\  }}|                     ¦   «         D ]\  }	}
||	v r|
|z  }
ŒŒ)|S )N©r    r   rh   ri   r  )rU   rž   r  )rf   r    r   rh   ri   r  r  r  rË   Úloss_keyr   s              r0   Úget_loss_dictz*EomtForUniversalSegmentation.get_loss_dict!  sŒ   € ð (,§~¢~Ø!5Ø!5Ø#Ø%Ø"7ð (6ñ (
ô (
ˆ	ð  Ô+×1Ò1Ñ3Ô3ð 	#ð 	#‰KˆC�Ø"+§/¢/Ñ"3Ô"3ð #ð #‘�˜$Ø˜(�?�?Ø˜F‘N�Døð#ð Ðr/   r  c                 óD   — t          |                     ¦   «         ¦  «        S rà   )rH   rx  )rf   r  s     r0   Úget_lossz%EomtForUniversalSegmentation.get_loss9  s   € Ý�9×#Ò#Ñ%Ô%Ñ&Ô&Ð&r/   Nr$   r=   c                 ó°  — d\  }}d}|€t          d¦  «        ‚|                      |¦  «        }	t          | j        ¦  «        D �]w\  }
}|
| j        | j        j        z
  k    ri| j        j        ddd…dd…f          	                    |	j
        d         dd¦  «                             |	j        ¦  «        }t          j        ||	fd¬¦  «        }	|
| j        | j        j        z
  k    �rË| j        s'| j        |
| j        z
  | j        j        z            dk    �r�|                      |	¦  «        }|                      |¦  «        \  }}||fz  }||fz  }t          j        |	j
        d         |	j
        d         |	j
        d         |	j        t          j        ¬¦  «        }t+          j        || j        d	¬
¦  «        }|                     |                     d¦  «        |                     d¦  «        d¦  «        }| j        j        }|| j        j        z   }|dk    |dd…d|…|d…f<   |                      || j        |
| j        z
  | j        j        z            |||j        ¬¦  «        }|dd…ddf          	                    d| j        j        dd¦  «        }|                     ¦   «                              | d¦  «        } ||	|¦  «        }	�Œy|                      |	¦  «        }|                      |¦  «        \  }}||fz  }||fz  }d}|�L|�Jd}tA          ||¦  «        D ]7\  }}|  !                    ||||d¬¦  «        }||  "                    |¦  «        z  }Œ8tG          |||||¬¦  «        S )ag  
        mask_labels (`list[torch.Tensor]`, *optional*):
            list of mask labels of shape `(num_labels, height, width)` to be fed to a model
        class_labels (`list[torch.LongTensor]`, *optional*):
            list of target class labels of shape `(num_labels, height, width)` to be fed to a model. They identify the
            labels of `mask_labels`, e.g. the label of `mask_labels[i][j]` if `class_labels[i][j]`.
        patch_offsets (`list[torch.Tensor]`, *optional*):
            list of tuples indicating the image index and start and end positions of patches for semantic segmentation.
        )r.   r.   Nz You have to specify pixel_valuesr   rC   r   r  )rl   rq   Úbilinear)ÚsizeÚmode)ÚprobÚnum_query_tokensÚencoder_start_tokensrl   .g    eÍÍÁrR  r  )r   r    r   r!   r$   )$re   r:  r°   r  r  r�   rá  rT  rË   rM  rP   rz   rl   r)   rÒ   rZ  r  r  Úpredictr¥   r½   rË  Úinterpolater  r  r"  rÖ   rJ  Ú_disable_attention_maskrd  rŽ   Úmasked_fillr¾   r  r  r   )rf   r7  rh   ri   r$   r=   Úmasks_queries_logits_per_layerÚclass_queries_logits_per_layerrU  r"   r×   Úlayer_modulerT  Únorm_hidden_statesr    r   Úinterpolated_logitsr%  r&  Úsequence_outputr   r  s                         r0   r�   z$EomtForUniversalSegmentation.forward<  sî  € ð* JPÑFÐ&Ð(FØˆàÐÝÐ?Ñ@Ô@Ð@àŸš¨Ñ5Ô5ˆå!*¨4¬;Ñ!7Ô!7ð .	Hñ .	HÑˆC�Ø�dÔ,¨t¬{Ô/EÑEÒEÐEØœ
Ô)¨$°°°°1°1°1¨*Ô5×<Ò<¸]Ô=PÐQRÔ=SÐUWÐY[Ñ\Ô\×_Ò_Ð`mÔ`tÑuÔu�Ý %¤	¨5°-Ð*@ÀaÐ HÑ HÔ H�à�dÔ,¨t¬{Ô/EÑEÒEÑEØ”ð FØ!%Ô!5°c¸DÔ<RÑ6RÐUYÔU`ÔUkÑ6kÔ!lÐopÒ!pÑ!pà%)§^¢^°MÑ%BÔ%BÐ"Ø=A¿\º\ÐJ\Ñ=]Ô=]Ñ:Ð$Ð&:à.Ð3GÐ2IÑIÐ.Ø.Ð3GÐ2IÑIÐ.å!&¤Ø!Ô'¨Ô*Ø!Ô'¨Ô*Ø!Ô'¨Ô*Ø(Ô/Ýœ*ð"ñ "ô "�õ '(¤mÐ4HÈtÌ~ÐdnÐ&oÑ&oÔ&oÐ#Ø&9×&>Ò&>Ø'×,Ò,¨QÑ/Ô/Ð1D×1IÒ1IÈ!Ñ1LÔ1LÈbñ'ô 'Ð#ð $(¤;Ô#:Ð Ø'7¸$¼/Ô:[Ñ'[Ð$ð ObÐdeÒNe�˜q˜q˜qÐ"3Ð#3Ð"3Ð5IÐ5JÐ5JÐJÑKð "&×!=Ò!=Ø"ØÔ-¨c°DÔ4JÑ.JÈTÌ[ÔMcÑ.cÔdØ%5Ø)=Ø)Ô0ð ">ñ "ô "�ð "0°°°°4¸°Ô!=×!DÒ!DÀRÈÌÔIhÐjlÐnpÑ!qÔ!q�Ø!/×!5Ò!5Ñ!7Ô!7×!CÒ!CÀ^ÀOÐUYÑ!ZÔ!Z�à(˜L¨¸ÑGÔGˆM‰MàŸ.š.¨Ñ7Ô7ˆà59·\²\À/Ñ5RÔ5RÑ2ÐÐ2Ø&Ð+?Ð*AÑAÐ&Ø&Ð+?Ð*AÑAÐ&àˆØÐ" |Ð'?ØˆDÝ>AØ.Ð0Nñ?ô ?ð 
1ð 
1Ñ:Ð$Ð&:ð !×.Ò.Ø)=Ø)=Ø +Ø!-Ø*.ð /ñ ô �	ð ˜Ÿš iÑ0Ô0Ñ0��å1ØØ!5Ø!5Ø-Ø'ð
ñ 
ô 
ð 	
r/   c                 ó   — | j         j        S rà   )r:  rF  r«  s    r0   Úget_input_embeddingsz1EomtForUniversalSegmentation.get_input_embeddings¦  s   € ØŒÔ/Ð/r/   râ   c                 ó¤  — |d d …d | j         j        …d d …f         }|                      |¦  «        }|d d …| j         j        | j        j        z   d …d d …f         }|                     dd¦  «        } |j        |j        d         dg| j        ¢R Ž }|  	                    |¦  «        }|  
                    |¦  «        }t          j        d||¦  «        }||fS )Nr   r5   r   rC   zbqc, bchw -> bqhw)r�   rÖ   r  r:  rJ  rÔ   rs  rP   r  r  r  r)   Úeinsum)rf   râ   Úquery_tokensÚclass_logitsÚprefix_tokensÚmask_logitss         r0   r'  z$EomtForUniversalSegmentation.predict©  sè   € Ø˜a˜a˜aÐ!: 4¤;Ô#:Ð!:¸A¸A¸AÐ=Ô>ˆØ×+Ò+¨LÑ9Ô9ˆà˜q˜q˜q $¤+Ô"9¸D¼OÔ<]Ñ"]Ð"_Ð"_ÐabÐabÐabÐbÔcˆØ%×/Ò/°°1Ñ5Ô5ˆà-˜Ô-¨mÔ.AÀ!Ô.DÀbÐZÈ4Ì>ÐZÐZÐZˆà—~’~ lÑ3Ô3ˆØ×*Ò*¨=Ñ9Ô9ˆå”lÐ#6¸ÀmÑTÔTˆà˜LÐ(Ð(r/   c                 ó†   — |dk     r:t          j        | j        d         ||¬¦  «        |k    }d| d d …d |…|d …f         |<   | S )Nr   r   rk   )r)   r{   rP   )Ú	attn_maskr$  r%  r&  rl   Úrandom_queriess         r0   r)  z4EomtForUniversalSegmentation._disable_attention_mask¹  sb   € à�!Š8ˆ8å"œZ¨	¬¸Ô(:Ð<LÐU[Ð\Ñ\Ô\Ð_cÒcˆNð VWˆI�a�a�aÐ*Ð*Ð*Ð,@Ð,AÐ,AÐAÔBÀ>ÑRàÐr/   )NNN)r%   r&   r'   r  r   rd   r   r"  r#  r  r  r   r   r   r-   r   r   r   r�   r2  r)   r'  Ústaticmethodr)  r‘   r’   s   @r0   r  r  ý  sÉ  ø€ € € € € ð %€Oð˜zð ð ð ð ð ð ð8à$ðð %ðð ð	ð
 ðð  $ C¨ KÔ0ðð 
ˆc�6ˆkÔ	ðð ð ð ð0' $ s¨F {Ô"3ð '¸ð 'ð 'ð 'ð 'ð  ØØð ,0Ø,0Ø-1ðe
ð e
àðe
ð ˜&”\ DÑ(ðe
ð ˜6”l TÑ)ð	e
ð
 ˜F”| dÑ*ðe
ð Ð+Ô,ðe
ð 
,ðe
ð e
ð e
ñ „^ñ „_ñ  Ôðe
ðN0ð 0ð 0ð)˜eœlð )ð )ð )ð )ð  ðð ñ „\ðð ð ð ð r/   r  )Fr­  )JÚcollections.abcr1  r÷  r   Údataclassesr   Únumpyr$  r)   Útorch.nn.functionalr   r9   rË  r   Ú r   rõ  Úactivationsr	   Ú
file_utilsr
   r   r   Úmodeling_layersr   Úmodeling_utilsr   r   Úprocessing_utilsr   Úutilsr   r   r   Úutils.genericr   Úutils.output_capturingr   Úconfiguration_eomtr   Úscipy.optimizer   Ú
accelerater   Úaccelerate.utilsr   r   r?   rK   rZ   r  r\   r�   r–   rš   rœ   r'  r<  rŽ   r_  ra  r{  r…  r”  r¡  r¯  r³  rÄ  rÏ  rÜ  ræ  rë  r  Ú__all__r.   r/   r0   ú<module>rO     sŸ  ðð* Ð Ð Ð Ø €€€Ø $Ð $Ð $Ð $Ð $Ð $Ø !Ð !Ð !Ð !Ð !Ð !à Ð Ð Ð Ø €€€Ø Ð Ð Ð Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø LÐ LÐ LÐ LÐ LÐ LÐ LÐ LÐ LÐ LØ 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø PÐ PÐ PÐ PÐ PÐ PÐ PÐ PÐ PÐ PØ 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø *Ð *Ð *Ð *Ð *Ð *ð ÐÑÔð 5Ø4Ð4Ð4Ð4Ð4Ð4àÐÑÔð (Ø'Ð'Ð'Ð'Ð'Ð'Ø'Ð'Ð'Ð'Ð'Ð'ð €ðð	ñ 	ô 	ð ð4ð 4ð 4ð 4ð 4¨ñ 4ô 4ñ „ñ	ô 	ð4ðB LQðð Ø”LðØ5:´\ðà
„\ðð ð ð ð@ ð °ð ¸6ð ð ð ð ð,°´ð ÀuÄ|ð ÐX]ÔXdð ð ð ð ð8gð gð gð gð g˜2œ9ñ gô gð gðT�fð  fð ¸ð Àð ð ð ð ð< u¤|ð ¸U¼\ð ÐVYð Ð^cÔ^jð ð ð ð ð(uð uð uð uð uˆrŒyñ uô uð uðp	ð ð ð ð ˜"œ)ñ ô ð ðB"ð "ð "ð "ð "�R”Yñ "ô "ð "ðX ð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð ð%ð ð%ð %ð %ð %ð.8)ð 8)ð 8)ð 8)ð 8)�B”Iñ 8)ô 8)ð 8)ðv+ð +ð +ð +ð +�R”Yñ +ô +ð +ðð ð ð ð ˆbŒiñ ô ð ð&(ð (ð (ð (ð (�B”Iñ (ô (ð (ð"%ð %ð %ð %ð %�2”9ñ %ô %ð %ð0'ð 'ð 'ð 'ð 'Ð*ñ 'ô 'ð 'ðTð ð ð ð �b”lñ ô ð ðð ð ð ð �R”Yñ ô ð ð2	ð 	ð 	ð 	ð 	�R”Yñ 	ô 	ð 	ðð ð ð ð �2”9ñ ô ð ð" ð-/ð -/ð -/ð -/ð -/˜/ñ -/ô -/ñ „ð-/ð` €ððñ ô ð
@ð @ð @ð @ð @Ð#6ñ @ô @ñô ð
@ðF !Ð"@Ð
A€€€r/   