Ë
    T^(hà  ã                   ó¶  — d Z ddlZddlmZ ddlmZmZmZmZ ddl	Z	ddl
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 ddlmZmZ ddlmZ ddlmZ  ej:                  e«      ZdZ e G d„ de«      «       Z! G d„ dejD                  «      Z# G d„ dejD                  «      Z$ G d„ dejD                  «      Z% G d„ dejD                  «      Z& G d„ dejD                  «      Z' G d„ dejD                  «      Z( G d„ dejD                  «      Z)dCd „Z* G d!„ d"ejD                  «      Z+ G d#„ d$ejD                  «      Z, G d%„ d&ejD                  «      Z-e	j\                  j^                  dDd'e0d(e1fd)„«       Z2 G d*„ d+ejD                  «      Z3 G d,„ d-ejD                  «      Z4 G d.„ d/ejD                  «      Z5 G d0„ d1ejD                  «      Z6 G d2„ d3ejD                  «      Z7 G d4„ d5ejD                  «      Z8 G d6„ d7ejD                  «      Z9 G d8„ d9ejD                  «      Z: G d:„ d;ejD                  «      Z; G d<„ d=e«      Z<d>Z=d?Z> ed@e=«       G dA„ dBe<«      «       Z?dBd=gZ@y)EzPyTorch ZoeDepth model.é    N)Ú	dataclass)ÚListÚOptionalÚTupleÚUnion)Únné   )ÚACT2FN)Úadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚreplace_return_docstrings)ÚDepthEstimatorOutput)ÚPreTrainedModel)ÚModelOutputÚlogging)Úload_backboneé   )ÚZoeDepthConfigr   c                   óî   — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eej                     ed<   dZeeej                  df      ed<   dZeeej                  df      ed<   y)	ÚZoeDepthDepthEstimatorOutputaâ  
    Extension of `DepthEstimatorOutput` to include domain logits (ZoeDepth specific).

    Args:
        loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
            Classification (or regression if config.num_labels==1) loss.
        predicted_depth (`torch.FloatTensor` of shape `(batch_size, height, width)`):
            Predicted depth for each pixel.

        domain_logits (`torch.FloatTensor` of shape `(batch_size, num_domains)`):
            Logits for each domain (e.g. NYU and KITTI) in case multiple metric heads are used.

        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, if the model has an embedding layer, +
            one for the output of each layer) of shape `(batch_size, num_channels, height, width)`.

            Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
        attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, patch_size,
            sequence_length)`.

            Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
            heads.
    NÚlossÚpredicted_depthÚdomain_logits.Úhidden_statesÚ
attentions)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   ÚtorchÚFloatTensorÚ__annotations__r   r   r   r   r   © ó    úl/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/zoedepth/modeling_zoedepth.pyr   r   ,   s†   … ñð2 )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø37€O�X˜e×/Ñ/Ñ0Ó7Ø15€M�8˜E×-Ñ-Ñ.Ó5Ø=A€M�8˜E %×"3Ñ"3°SÐ"8Ñ9Ñ:ÓAØ:>€J�˜˜u×0Ñ0°#Ð5Ñ6Ñ7Ô>r$   r   c                   óf   ‡ — e Zd ZdZˆ fd„Zdeej                     deej                     fd„Zˆ xZ	S )ÚZoeDepthReassembleStageaE  
    This class reassembles the hidden states of the backbone into image-like feature representations at various
    resolutions.

    This happens in 3 stages:
    1. Map the N + 1 tokens to a set of N tokens, by taking into account the readout ([CLS]) token according to
       `config.readout_type`.
    2. Project the channel dimension of the hidden states according to `config.neck_hidden_sizes`.
    3. Resizing the spatial dimensions (height, width).

    Args:
        config (`[ZoeDepthConfig]`):
            Model configuration class defining the model architecture.
    c           	      óN  •— t         ‰| �  «        |j                  | _        t        j                  «       | _        t        |j                  |j                  «      D ],  \  }}| j
                  j                  t        |||¬«      «       Œ. |j                  dk(  rŽt        j                  «       | _        |j                  }|j                  D ]Y  }| j                  j                  t        j                  t        j                  d|z  |«      t        |j                      «      «       Œ[ y y )N)ÚchannelsÚfactorÚprojecté   )ÚsuperÚ__init__Úreadout_typer   Ú
ModuleListÚlayersÚzipÚneck_hidden_sizesÚreassemble_factorsÚappendÚZoeDepthReassembleLayerÚreadout_projectsÚbackbone_hidden_sizeÚ
SequentialÚLinearr
   Ú
hidden_act)ÚselfÚconfigÚneck_hidden_sizer*   Úhidden_sizeÚ_Ú	__class__s         €r%   r.   z ZoeDepthReassembleStage.__init__^   sñ   ø€ Ü‰ÑÔà"×/Ñ/ˆÔÜ—m‘m“oˆŒä(+¨F×,DÑ,DÀf×F_ÑF_Ó(`ò 	jÑ$Ð˜fØ�K‰K×ÑÔ6°vÐHXÐagÔhÕið	jð ×Ñ )Ò+Ü$&§M¡M£OˆDÔ!Ø ×5Ñ5ˆKØ×-Ñ-ò �Ø×%Ñ%×,Ñ,Ü—M‘M¤"§)¡)¨A°©O¸[Ó"IÌ6ÐRX×RcÑRcÑKdÓeõñð ,r$   r   Úreturnc                 óN  — |d   j                   d   }t        j                  |d¬«      }|dd…df   |dd…dd…f   }}|j                   \  }}}|j                  ||||«      }|j	                  dddd«      j                  «       }| j                  dk(  rZ|j                  d«      j	                  d«      }|j                  d¬«      j                  |«      }	t        j                  ||	fd	«      }n#| j                  d
k(  r||j                  d	«      z   }g }
t        |j                  |d¬«      «      D ]t  \  }}| j                  dk(  r | j                  |   |«      }|j	                  ddd«      j                  |d	||«      } | j                  |   |«      }|
j                  |«       Œv |
S )zÇ
        Args:
            hidden_states (`List[torch.FloatTensor]`, each of shape `(batch_size, sequence_length + 1, hidden_size)`):
                List of hidden states from the backbone.
        r   ©ÚdimNr   r	   r,   r+   )r   r,   r   éÿÿÿÿÚadd)Úshaper    ÚcatÚreshapeÚpermuteÚ
contiguousr/   ÚflattenÚ	unsqueezeÚ	expand_asÚ	enumerateÚsplitr7   r1   r5   )r<   r   Úpatch_heightÚpatch_widthÚ
batch_sizeÚ	cls_tokenÚtotal_batch_sizeÚsequence_lengthÚnum_channelsÚreadoutÚoutÚ	stage_idxÚhidden_states                r%   ÚforwardzZoeDepthReassembleStage.forwardo   s¸  € ð # 1Ñ%×+Ñ+¨AÑ.ˆ
ô Ÿ	™	 -°QÔ7ˆà#0²°A°Ñ#6¸ÂaÈÉÀeÑ8L�=ˆ	à:G×:MÑ:MÑ7Ð˜/¨<Ø%×-Ñ-Ð.>ÀÈkÐ[gÓhˆØ%×-Ñ-¨a°°A°qÓ9×DÑDÓFˆà×Ñ 	Ò)à)×1Ñ1°!Ó4×<Ñ<¸YÓGˆMØ×)Ñ)¨aÐ)Ó0×:Ñ:¸=ÓIˆGô "ŸI™I }°gÐ&>ÀÓC‰MØ×Ñ %Ò'Ø)¨I×,?Ñ,?ÀÓ,CÑCˆMàˆÜ'0°×1DÑ1DÀZÐUVÐ1DÓ1WÓ'Xò 	%Ñ#ˆI�|Ø× Ñ  IÒ-Ø?˜t×4Ñ4°YÑ?ÀÓM�ð (×/Ñ/°°1°aÓ8×@Ñ@ÀÈRÐQ]Ð_jÓkˆLØ1˜4Ÿ;™; yÑ1°,Ó?ˆLØ�J‰J�|Õ$ð	%ð ˆ
r$   ©
r   r   r   r   r.   r   r    ÚTensorr]   Ú__classcell__©rA   s   @r%   r'   r'   N   s6   ø„ ñôð"& T¨%¯,©,Ñ%7ð &ÐW[Ð\a×\hÑ\hÑWi÷ &r$   r'   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )r6   c           	      ó^  •— t         ‰| �  «        |j                  }t        j                  ||d¬«      | _        |dkD  r t        j                  ||||d¬«      | _        y |dk(  rt        j                  «       | _        y |dk  r,t        j                  ||dt        d|z  «      d¬«      | _        y y )Nr   )Úin_channelsÚout_channelsÚkernel_sizer   ©rf   ÚstrideÚpaddingr	   )
r-   r.   r8   r   ÚConv2dÚ
projectionÚConvTranspose2dÚresizeÚIdentityÚint)r<   r=   r)   r*   r?   rA   s        €r%   r.   z ZoeDepthReassembleLayer.__init__™   s—   ø€ Ü‰ÑÔà×1Ñ1ˆÜŸ)™)°È(Ð`aÔbˆŒð �AŠ:Ü×,Ñ,¨X°xÈVÐ\bÐlmÔnˆD�KØ�qŠ[ÜŸ+™+›-ˆD�KØ�aŠZäŸ)™) H¨hÀAÌcÐRSÐV\ÑR\ËoÐghÔiˆD�Kð r$   c                 óJ   — | j                  |«      }| j                  |«      }|S ©N)rk   rm   ©r<   r\   s     r%   r]   zZoeDepthReassembleLayer.forward©   s$   € Ø—‘ |Ó4ˆØ—{‘{ <Ó0ˆØÐr$   ©r   r   r   r.   r]   r`   ra   s   @r%   r6   r6   ˜   s   ø„ ôjö r$   r6   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚZoeDepthFeatureFusionStagec                 óâ   •— t         ‰| �  «        t        j                  «       | _        t        t        |j                  «      «      D ]&  }| j                  j                  t        |«      «       Œ( y rq   )
r-   r.   r   r0   r1   ÚrangeÚlenr3   r5   ÚZoeDepthFeatureFusionLayer)r<   r=   r@   rA   s      €r%   r.   z#ZoeDepthFeatureFusionStage.__init__±   sT   ø€ Ü‰ÑÔÜ—m‘m“oˆŒÜ”s˜6×3Ñ3Ó4Ó5ò 	CˆAØ�K‰K×ÑÔ9¸&ÓAÕBñ	Cr$   c                 ó¤   — |d d d…   }g }d }t        || j                  «      D ]*  \  }}|€	 ||«      }n	 |||«      }|j                  |«       Œ, |S )NrF   )r2   r1   r5   )r<   r   Úfused_hidden_statesÚfused_hidden_stater\   Úlayers         r%   r]   z"ZoeDepthFeatureFusionStage.forward·   sq   € à%¡d¨ dÑ+ˆà ÐØ!ÐÜ#& }°d·k±kÓ#Bò 	;ÑˆL˜%Ø!Ð)á%*¨<Ó%8Ñ"á%*Ð+=¸|Ó%LÐ"Ø×&Ñ&Ð'9Õ:ð	;ð #Ð"r$   rs   ra   s   @r%   ru   ru   °   s   ø„ ôCö#r$   ru   c                   óZ   ‡ — e Zd ZdZˆ fd„Zdej                  dej                  fd„Zˆ xZS )ÚZoeDepthPreActResidualLayerz®
    ResidualConvUnit, pre-activate residual unit.

    Args:
        config (`[ZoeDepthConfig]`):
            Model configuration class defining the model architecture.
    c                 óœ  •— t         ‰| �  «        |j                  | _        |j                  �|j                  n| j                   }t        j                  «       | _        t        j                  |j                  |j                  ddd|¬«      | _
        t        j                  «       | _        t        j                  |j                  |j                  ddd|¬«      | _        | j                  rat        j                  |j                  |j                  ¬«      | _        t        j                  |j                  |j                  ¬«      | _        y y )Nr	   r   )rf   rh   ri   Úbias)Úeps)r-   r.   Ú!use_batch_norm_in_fusion_residualÚuse_batch_normÚuse_bias_in_fusion_residualr   ÚReLUÚactivation1rj   Úfusion_hidden_sizeÚconvolution1Úactivation2Úconvolution2ÚBatchNorm2dÚbatch_norm_epsÚbatch_norm1Úbatch_norm2)r<   r=   r…   rA   s      €r%   r.   z$ZoeDepthPreActResidualLayer.__init__Ó   s  ø€ Ü‰ÑÔà$×FÑFˆÔð ×1Ñ1Ð=ð ×.Ò.à×(Ñ(Ð(ð 	$ô Ÿ7™7›9ˆÔÜŸI™IØ×%Ñ%Ø×%Ñ%ØØØØ,ô
ˆÔô Ÿ7™7›9ˆÔÜŸI™IØ×%Ñ%Ø×%Ñ%ØØØØ,ô
ˆÔð ×ÒÜ!Ÿ~™~¨f×.GÑ.GÈV×MbÑMbÔcˆDÔÜ!Ÿ~™~¨f×.GÑ.GÈV×MbÑMbÔcˆDÕð r$   r\   rB   c                 ó  — |}| j                  |«      }| j                  |«      }| j                  r| j                  |«      }| j	                  |«      }| j                  |«      }| j                  r| j                  |«      }||z   S rq   )r‡   r‰   r„   rŽ   rŠ   r‹   r�   ©r<   r\   Úresiduals      r%   r]   z#ZoeDepthPreActResidualLayer.forwardõ   s„   € ØˆØ×'Ñ'¨Ó5ˆà×(Ñ(¨Ó6ˆà×ÒØ×+Ñ+¨LÓ9ˆLà×'Ñ'¨Ó5ˆØ×(Ñ(¨Ó6ˆà×ÒØ×+Ñ+¨LÓ9ˆLà˜hÑ&Ð&r$   )	r   r   r   r   r.   r    r_   r]   r`   ra   s   @r%   r   r   É   s*   ø„ ñô dðD' E§L¡Lð '°U·\±\÷ 'r$   r   c                   ó,   ‡ — e Zd ZdZdˆ fd„	Zdd„Zˆ xZS )ry   a8  Feature fusion layer, merges feature maps from different stages.

    Args:
        config (`[ZoeDepthConfig]`):
            Model configuration class defining the model architecture.
        align_corners (`bool`, *optional*, defaults to `True`):
            The align_corner setting for bilinear upsample.
    c                 óÔ   •— t         ‰| �  «        || _        t        j                  |j
                  |j
                  dd¬«      | _        t        |«      | _        t        |«      | _	        y )Nr   T)rf   r�   )
r-   r.   Úalign_cornersr   rj   rˆ   rk   r   Úresidual_layer1Úresidual_layer2)r<   r=   r•   rA   s      €r%   r.   z#ZoeDepthFeatureFusionLayer.__init__  sT   ø€ Ü‰ÑÔà*ˆÔäŸ)™) F×$=Ñ$=¸v×?XÑ?XÐfgÐnrÔsˆŒä:¸6ÓBˆÔÜ:¸6ÓBˆÕr$   c                 ó€  — |�l|j                   |j                   k7  r?t        j                  j                  ||j                   d   |j                   d   fdd¬«      }|| j	                  |«      z   }| j                  |«      }t        j                  j                  |dd| j                  ¬«      }| j                  |«      }|S )Nr,   r	   ÚbilinearF©ÚsizeÚmoder•   ©Úscale_factorrœ   r•   )rH   r   Ú
functionalÚinterpolater–   r—   r•   rk   r‘   s      r%   r]   z"ZoeDepthFeatureFusionLayer.forward  s¾   € ØÐØ×!Ñ! X§^¡^Ò3ÜŸ=™=×4Ñ4Ø L×$6Ñ$6°qÑ$9¸<×;MÑ;MÈaÑ;PÐ#QÐXbÐrwð 5ó �ð (¨$×*>Ñ*>¸xÓ*HÑHˆLà×+Ñ+¨LÓ9ˆÜ—}‘}×0Ñ0Ø q¨zÈ×I[ÑI[ð 1ó 
ˆð —‘ |Ó4ˆàÐr$   )Trq   ©r   r   r   r   r.   r]   r`   ra   s   @r%   ry   ry     s   ø„ ñõC÷r$   ry   c                   óf   ‡ — e Zd ZdZˆ fd„Zdeej                     deej                     fd„Zˆ xZ	S )ÚZoeDepthNeckaO  
    ZoeDepthNeck. A neck is a module that is normally used between the backbone and the head. It takes a list of tensors as
    input and produces another list of tensors as output. For ZoeDepth, it includes 2 stages:

    * ZoeDepthReassembleStage
    * ZoeDepthFeatureFusionStage.

    Args:
        config (dict): config dict.
    c           
      ó–  •— t         ‰| �  «        || _        |j                  � |j                  j                  dv rd | _        nt        |«      | _        t        j                  «       | _	        |j                  D ]?  }| j                  j                  t        j                  ||j                  ddd¬«      «       ŒA t        |«      | _        y )N)Úswinv2r	   r   F)rf   ri   r�   )r-   r.   r=   Úbackbone_configÚ
model_typeÚreassemble_stager'   r   r0   Úconvsr3   r5   rj   rˆ   ru   Úfusion_stage)r<   r=   ÚchannelrA   s      €r%   r.   zZoeDepthNeck.__init__:  s«   ø€ Ü‰ÑÔØˆŒð ×!Ñ!Ð-°&×2HÑ2H×2SÑ2SÐWaÑ2aØ$(ˆDÕ!ä$;¸FÓ$CˆDÔ!ä—]‘]“_ˆŒ
Ø×/Ñ/ò 	sˆGØ�J‰J×ÑœbŸi™i¨°×1JÑ1JÐXYÐcdÐkpÔqÕrð	sô 7°vÓ>ˆÕr$   r   rB   c                 óŠ  — t        |t        t        f«      st        d«      ‚t	        |«      t	        | j
                  j                  «      k7  rt        d«      ‚| j                  �| j                  |||«      }t        |«      D ��cg c]  \  }} | j                  |   |«      ‘Œ }}}| j                  |«      }||d   fS c c}}w )zñ
        Args:
            hidden_states (`List[torch.FloatTensor]`, each of shape `(batch_size, sequence_length, hidden_size)` or `(batch_size, hidden_size, height, width)`):
                List of hidden states from the backbone.
        z2hidden_states should be a tuple or list of tensorszOThe number of hidden states should be equal to the number of neck hidden sizes.rF   )Ú
isinstanceÚtupleÚlistÚ	TypeErrorrx   r=   r3   Ú
ValueErrorr¨   rP   r©   rª   )r<   r   rR   rS   ÚiÚfeatureÚfeaturesÚoutputs           r%   r]   zZoeDepthNeck.forwardK  sº   € ô ˜-¬%´¨Ô7ÜÐPÓQÐQäˆ}Ó¤ T§[¡[×%BÑ%BÓ!CÒCÜÐnÓoÐoð × Ñ Ð,Ø ×1Ñ1°-ÀÈ{Ó[ˆMä=FÀ}Ó=U×V©z¨q°'�M�D—J‘J˜q‘M 'Õ*ÐVˆÑVð ×"Ñ" 8Ó,ˆà�x ‘|Ð#Ð#ùó Ws   ÂB?r^   ra   s   @r%   r£   r£   -  s6   ø„ ñ	ô?ð"$ T¨%¯,©,Ñ%7ð $ÐW[Ð\a×\hÑ\hÑWi÷ $r$   r£   c                   ó`   ‡ — e Zd ZdZˆ fd„Zdeej                     dej                  fd„Zˆ xZ	S )Ú#ZoeDepthRelativeDepthEstimationHeada  
    Relative depth estimation head consisting of 3 convolutional layers. It progressively halves the feature dimension and upsamples
    the predictions to the input resolution after the first convolutional layer (details can be found in DPT's paper's
    supplementary material).
    c                 óè  •— t         ‰| �  «        |j                  | _        d | _        |j                  rt        j                  ddddd¬«      | _        |j                  }t        j                  ||dz  ddd¬«      | _        t        j                  ddd	¬
«      | _
        t        j                  |dz  |j                  ddd¬«      | _        t        j                  |j                  dddd¬«      | _        y )Né   )r	   r	   )r   r   rg   r,   r	   r   r™   Tr�   r   )r-   r.   Úhead_in_indexrk   Úadd_projectionr   rj   rˆ   Úconv1ÚUpsampleÚupsampleÚnum_relative_featuresÚconv2Úconv3)r<   r=   r´   rA   s      €r%   r.   z,ZoeDepthRelativeDepthEstimationHead.__init__j  sÇ   ø€ Ü‰ÑÔà#×1Ñ1ˆÔàˆŒØ× Ò Ü Ÿi™i¨¨S¸fÈVÐ]cÔdˆDŒOà×,Ñ,ˆÜ—Y‘Y˜x¨°Q©ÀAÈaÐYZÔ[ˆŒ
ÜŸ™°¸ÐSWÔXˆŒÜ—Y‘Y˜x¨1™}¨f×.JÑ.JÐXYÐbcÐmnÔoˆŒ
Ü—Y‘Y˜v×;Ñ;¸QÈAÐVWÐabÔcˆ�
r$   r   rB   c                 ó®  — || j                      }| j                  �+| j                  |«      } t        j                  «       |«      }| j	                  |«      }| j                  |«      }| j                  |«      } t        j                  «       |«      }|}| j                  |«      } t        j                  «       |«      }|j                  d¬«      }||fS )Nr   rD   )	rº   rk   r   r†   r¼   r¾   rÀ   rÁ   Úsqueeze)r<   r   r´   r   s       r%   r]   z+ZoeDepthRelativeDepthEstimationHead.forwardy  s»   € à% d×&8Ñ&8Ñ9ˆà�?‰?Ð&Ø ŸO™O¨MÓ:ˆMØ%œBŸG™G›I mÓ4ˆMàŸ
™
 =Ó1ˆØŸ™ mÓ4ˆØŸ
™
 =Ó1ˆØ!œŸ™›	 -Ó0ˆà ˆØŸ
™
 =Ó1ˆØ!œŸ™›	 -Ó0ˆà'×/Ñ/°AÐ/Ó6ˆà Ð(Ð(r$   r^   ra   s   @r%   r·   r·   c  s.   ø„ ñôdð) T¨%¯,©,Ñ%7ð )¸E¿L¹L÷ )r$   r·   c                 ó¼   — | |z   } ||z   }| t        j                  | «      z  |t        j                  |«      z  z
  | |z
  t        j                  | |z
  |z   «      z  z
  S )z%log(nCk) using stirling approximation)r    Úlog)ÚnÚkr‚   s      r%   Ú	log_binomrÈ   �  sX   € à	ˆC‰€AØ	ˆC‰€AØŒu�y‰y˜‹|Ñ˜a¤%§)¡)¨A£,Ñ.Ñ.°!°a±%¼5¿9¹9ÀQÈÁUÈSÁ[Ó;QÑ1QÑQÐQr$   c                   ó@   ‡ — e Zd Zdej                  fˆ fd„	Zdd„Zˆ xZS )ÚLogBinomialSoftmaxr¹   c           	      ó@  •— t         ‰| �  «        || _        || _        | j	                  dt        j                  d|«      j                  dddd«      d¬«       | j	                  dt        j                  | j                  dz
  g«      j                  dddd«      d¬«       y)	a7  Compute log binomial distribution for n_classes

        Args:
            n_classes (`int`, *optional*, defaults to 256):
                Number of output classes.
            act (`torch.nn.Module`, *optional*, defaults to `torch.softmax`):
                Activation function to apply to the output.
        Úk_idxr   r   rF   F)Ú
persistentÚ	k_minus_1N)	r-   r.   rÇ   ÚactÚregister_bufferr    ÚarangeÚviewÚtensor)r<   Ú	n_classesrÏ   rA   s      €r%   r.   zLogBinomialSoftmax.__init__—  s�   ø€ ô 	‰ÑÔØˆŒØˆŒØ×Ñ˜W¤e§l¡l°1°iÓ&@×&EÑ&EÀaÈÈQÐPQÓ&RÐ_dÐÔeØ×Ñ˜[¬%¯,©,¸¿¹À¹
°|Ó*D×*IÑ*IÈ!ÈRÐQRÐTUÓ*VÐchÐÕir$   c                 ó¶  — |j                   dk(  r|j                  d«      }t        j                  d|z
  |d«      }t        j                  ||d«      }t	        | j
                  | j                  «      | j                  t        j                  |«      z  z   | j
                  | j                  z
  t        j                  |«      z  z   }| j                  ||z  d¬«      S )a°  Compute the log binomial distribution for probabilities.

        Args:
            probabilities (`torch.Tensor` of shape `(batch_size, num_channels, height, width)`):
                Tensor containing probabilities of each class.
            temperature (`float` or `torch.Tensor` of shape `(batch_size, num_channels, height, width)`, *optional*, defaults to 1):
                Temperature of distribution.
            eps (`float`, *optional*, defaults to 1e-4):
                Small number for numerical stability.

        Returns:
            `torch.Tensor` of shape `(batch_size, num_channels, height, width)`:
                Log binomial distribution logbinomial(p;t).
        r	   r   rD   )	ÚndimrN   r    ÚclamprÈ   rÎ   rÌ   rÅ   rÏ   )r<   ÚprobabilitiesÚtemperaturer‚   Úone_minus_probabilitiesÚys         r%   r]   zLogBinomialSoftmax.forward¦  s»   € ð ×Ñ Ò"Ø)×3Ñ3°AÓ6ˆMä"'§+¡+¨a°-Ñ.?ÀÀaÓ"HÐÜŸ™ M°3¸Ó:ˆä�d—n‘n d§j¡jÓ1Ø�j‰jœ5Ÿ9™9 ]Ó3Ñ3ñ4à�~‰~ §
¡
Ñ*¬e¯i©iÐ8OÓ.PÑPñQð 	
ð
 �x‰x˜˜K™¨QˆxÓ/Ð/r$   )ç      ð?ç-Cëâ6?)r   r   r   r    Úsoftmaxr.   r]   r`   ra   s   @r%   rÊ   rÊ   –  s   ø„ Ø!$¨%¯-©-õ j÷0r$   rÊ   c                   ó*   ‡ — e Zd Z	 	 dˆ fd„	Zd„ Zˆ xZS )Ú%ZoeDepthConditionalLogBinomialSoftmaxc                 ó¬  •— t         ‰| �  «        ||z   |z  }t        j                  t        j                  ||z   |ddd¬«      t        j
                  «       t        j                  |dddd¬«      t        j                  «       «      | _        d| _        |j                  | _	        |j                  | _
        t        |t        j                  ¬«      | _        y)aß  Per-pixel MLP followed by a Conditional Log Binomial softmax.

        Args:
            in_features (`int`):
                Number of input channels in the main feature.
            condition_dim (`int`):
                Number of input channels in the condition feature.
            n_classes (`int`, *optional*, defaults to 256):
                Number of classes.
            bottleneck_factor (`int`, *optional*, defaults to 2):
                Hidden dim factor.

        r   r   rg   é   rÝ   )rÏ   N)r-   r.   r   r9   rj   ÚGELUÚSoftplusÚmlpÚp_epsÚmax_tempÚmin_temprÊ   r    rÞ   Úlog_binomial_transform)r<   r=   Úin_featuresÚcondition_dimrÔ   Úbottleneck_factorÚ
bottleneckrA   s          €r%   r.   z.ZoeDepthConditionalLogBinomialSoftmax.__init__Ã  s£   ø€ ô* 	‰ÑÔà! MÑ1Ð6GÑGˆ
Ü—=‘=Ü�I‰I�k MÑ1°:È1ÐUVÐ`aÔbÜ�G‰G‹Iä�I‰I�j %°Q¸qÈ!ÔLÜ�K‰K‹Mó
ˆŒð ˆŒ
ØŸ™ˆŒØŸ™ˆŒÜ&8¸ÌÏÉÔ&VˆÕ#r$   c                 óÖ  — | j                  t        j                  ||fd¬«      «      }|dd…dd…df   |dd…dd…df   }}|| j                  z   }|dd…ddf   |dd…ddf   |dd…ddf   z   z  }|| j                  z   }|dd…ddf   |dd…ddf   |dd…ddf   z   z  }|j	                  d«      }| j
                  | j                  z
  |z  | j                  z   }| j                  ||«      S )az  
        Args:
            main_feature (`torch.Tensor` of shape `(batch_size, num_channels, height, width)`):
                Main feature.
            condition_feature (torch.Tensor of shape `(batch_size, num_channels, height, width)`):
                Condition feature.

        Returns:
            `torch.Tensor`:
                Output log binomial distribution
        r   rD   Nr,   .r   )rå   r    Úconcatræ   rN   rç   rè   ré   )r<   Úmain_featureÚcondition_featureÚprobabilities_and_temperaturerØ   rÙ   s         r%   r]   z-ZoeDepthConditionalLogBinomialSoftmax.forwardè  s  € ð )-¯©´·±¸|ÐM^Ð>_ÐefÔ1gÓ(hÐ%à)ª!¨R¨a¨R°¨*Ñ5Ø)ª!¨Q©R°¨*Ñ5ð #ˆð
 &¨¯
©
Ñ2ˆØ%¢a¨¨C iÑ0°MÂ!ÀQÈÀ)Ñ4LÈ}Ò]^Ð`aÐcfÐ]fÑOgÑ4gÑhˆà! D§J¡JÑ.ˆØ!¢! Q¨ )Ñ,°ºA¸qÀ#¸IÑ0FÈÒUVÐXYÐ[^ÐU^ÑI_Ñ0_Ñ`ˆØ!×+Ñ+¨AÓ.ˆØ—}‘} t§}¡}Ñ4¸ÑCÀdÇmÁmÑSˆà×*Ñ*¨=¸+ÓFÐFr$   )r¹   r,   rs   ra   s   @r%   rà   rà   Â  s   ø„ ð Øõ#WöJGr$   rà   c                   ó&   ‡ — e Zd Zdˆ fd„	Zd„ Zˆ xZS )ÚZoeDepthSeedBinRegressorc                 óÌ  •— t         ‰| �  «        |j                  | _        |j                  | _        || _        || _        t        j                  | j                  |ddd«      | _	        t        j                  d¬«      | _        t        j                  ||ddd«      | _        | j                  dk(  rt        j                  d¬«      | _        yt        j                  «       | _        y)ad  Bin center regressor network.

        Can be "normed" or "unnormed". If "normed", bin centers are bounded on the (min_depth, max_depth) interval.

        Args:
            config (`int`):
                Model configuration.
            n_bins (`int`, *optional*, defaults to 16):
                Number of bin centers.
            mlp_dim (`int`, *optional*, defaults to 256):
                Hidden dimension.
            min_depth (`float`, *optional*, defaults to 1e-3):
                Min depth value.
            max_depth (`float`, *optional*, defaults to 10):
                Max depth value.
        r   r   T©ÚinplaceÚnormedN)r-   r.   Úbottleneck_featuresrê   Úbin_centers_typeÚ	min_depthÚ	max_depthr   rj   r¼   r†   Úact1rÀ   rä   Úact2)r<   r=   Ún_binsÚmlp_dimrû   rü   rA   s         €r%   r.   z!ZoeDepthSeedBinRegressor.__init__  s­   ø€ ô" 	‰ÑÔà!×5Ñ5ˆÔØ &× 7Ñ 7ˆÔØ"ˆŒØ"ˆŒä—Y‘Y˜t×/Ñ/°¸!¸QÀÓBˆŒ
Ü—G‘G DÔ)ˆŒ	Ü—Y‘Y˜w¨°°1°aÓ8ˆŒ
Ø-1×-BÑ-BÀhÒ-N”B—G‘G DÔ)ˆ�	ÔTV×T_ÑT_ÓTaˆ�	r$   c                 óæ  — | j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j                  dk(  r›|dz   }||j                  dd¬«      z  }| j                  | j                  z
  |z  }t        j                  j                  |dd| j                  ¬«      }t        j                  |d¬	«      }d
|dd…dd…df   |dd…dd…df   z   z  }||fS ||fS )z]
        Returns tensor of bin_width vectors (centers). One vector b for every pixel
        rø   çü©ñÒMbP?r   T©rE   Úkeepdim)r   r   r   r   r   r   Úconstant)rœ   ÚvaluerD   g      à?NrF   .)r¼   rý   rÀ   rþ   rú   Úsumrü   rû   r   rŸ   Úpadr    Úcumsum)r<   ÚxÚbin_centersÚbin_widths_normedÚ
bin_widthsÚ	bin_edgess         r%   r]   z ZoeDepthSeedBinRegressor.forward#  sù   € ð �J‰J�q‹MˆØ�I‰I�a‹LˆØ�J‰J�q‹MˆØ—i‘i “lˆà× Ñ  HÒ,Ø%¨Ñ,ˆKØ +¨k¯o©oÀ!ÈT¨oÓ.RÑ RÐàŸ.™.¨4¯>©>Ñ9Ð=NÑNˆJäŸ™×*Ñ*¨:Ð7IÐPZÐbf×bpÑbpÐ*ÓqˆJäŸ™ Z°QÔ7ˆIà ª1¨c¨r¨c°3¨;Ñ!7¸)ÂAÀqÁrÈ3ÀJÑ:OÑ!OÑPˆKØ$ kÐ1Ð1ð  Ð+Ð+r$   )é   r¹   r  é
   rs   ra   s   @r%   rô   rô     s   ø„ õbö:,r$   rô   ÚalphaÚgammac                 óN   — | j                  d|| j                  |«      z  z   «      S )a:  Inverse attractor: dc = dx / (1 + alpha*dx^gamma), where dx = a - c, a = attractor point, c = bin center, dc = shift in bin center
    This is the default one according to the accompanying paper.

    Args:
        dx (`torch.Tensor`):
            The difference tensor dx = Ai - Cj, where Ai is the attractor point and Cj is the bin center.
        alpha (`float`, *optional*, defaults to 300):
            Proportional Attractor strength. Determines the absolute strength. Lower alpha = greater attraction.
        gamma (`int`, *optional*, defaults to 2):
            Exponential Attractor strength. Determines the "region of influence" and indirectly number of bin centers affected.
            Lower gamma = farther reach.

    Returns:
        torch.Tensor: Delta shifts - dc; New bin centers = Old bin centers + dc
    r   )ÚdivÚpow)Údxr  r  s      r%   Úinv_attractorr  =  s%   € ð" �6‰6�!�e˜bŸf™f U›mÑ+Ñ+Ó,Ð,r$   c                   ó0   ‡ — e Zd Z	 	 	 	 dˆ fd„	Zdd„Zˆ xZS )ÚZoeDepthAttractorLayerc                 óÔ  •— t         ‰	| �  «        |j                  | _        |j                  | _        |j                  | _        || _        || _	        || _
        || _        || _        |j                  x}}t        j                  ||ddd«      | _        t        j"                  d¬«      | _        t        j                  ||dz  ddd«      | _        t        j"                  d¬«      | _        y)zq
        Attractor layer for bin centers. Bin centers are bounded on the interval (min_depth, max_depth)
        r   r   Trö   r,   N)r-   r.   Úattractor_alphar  Úattractor_gammaÚgemmaÚattractor_kindÚkindÚn_attractorsrÿ   rû   rü   Úmemory_efficientÚbin_embedding_dimr   rj   r¼   r†   rý   rÀ   rþ   ©
r<   r=   rÿ   r   rû   rü   r!  rê   r   rA   s
            €r%   r.   zZoeDepthAttractorLayer.__init__R  sÃ   ø€ ô 	‰ÑÔà×+Ñ+ˆŒ
Ø×+Ñ+ˆŒ
Ø×)Ñ)ˆŒ	à(ˆÔØˆŒØ"ˆŒØ"ˆŒØ 0ˆÔð !'× 8Ñ 8Ð8ˆ�gÜ—Y‘Y˜{¨G°Q¸¸1Ó=ˆŒ
Ü—G‘G DÔ)ˆŒ	Ü—Y‘Y˜w¨°qÑ(8¸!¸QÀÓBˆŒ
Ü—G‘G DÔ)ˆ�	r$   c                 ó˜  — |�7|r0t         j                  j                  ||j                  dd dd¬«      }||z   }| j	                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }|dz   }|j                  \  }}}}	|j                  || j                  d||	«      }|dd…dd…dd	f   }
t         j                  j                  |||	fdd¬«      }| j                  sct        j                  t        j                  d
œ| j                     } |t        |
j!                  d«      |j!                  d«      z
  «      d¬«      }n�t        j"                  ||j$                  ¬«      }t'        | j                  «      D ]*  }|t        |
dd…|d	f   j!                  d«      |z
  «      z  }Œ, | j                  dk(  r|| j                  z  }||z   }| j(                  | j*                  z
  |z  | j*                  z   }t        j,                  |d¬«      \  }}t        j.                  || j*                  | j(                  «      }||fS )ao  
        The forward pass of the attractor layer. This layer predicts the new bin centers based on the previous bin centers
        and the attractor points (the latter are predicted by the MLP).

        Args:
            x (`torch.Tensor` of shape `(batch_size, num_channels, height, width)`):
                Feature block.
            prev_bin (`torch.Tensor` of shape `(batch_size, prev_number_of_bins, height, width)`):
                Previous bin centers normed.
            prev_bin_embedding (`torch.Tensor`, *optional*):
                Optional previous bin embeddings.
            interpolate (`bool`, *optional*, defaults to `True`):
                Whether to interpolate the previous bin embeddings to the size of the input features.

        Returns:
            `Tuple[`torch.Tensor`, `torch.Tensor`]:
                New bin centers normed and scaled.
        Néþÿÿÿr™   T©rœ   r•   r  r,   r   .©Úmeanr  r   rD   ©Údevicer(  )r   rŸ   r    rH   r¼   rý   rÀ   rþ   rÒ   r   r!  r    r(  r  r  r  rN   Ú
zeros_liker*  rw   rü   rû   ÚsortÚclip)r<   r
  Úprev_binÚprev_bin_embeddingr    Ú
attractorsrT   r@   ÚheightÚwidthÚattractors_normedr  ÚfuncÚdelta_cr²   Úbin_new_centerss                   r%   r]   zZoeDepthAttractorLayer.forwardq  s%  € ð& Ð)ÙÜ%'§]¡]×%>Ñ%>Ø&¨¯©°°¨¸:ÐUYð &?ó &Ð"ð Ð&Ñ&ˆAà�J‰J�q‹MˆØ�I‰I�a‹LˆØ�J‰J�q‹MˆØ—Y‘Y˜q“\ˆ
à $Ñ&ˆ
Ø'1×'7Ñ'7Ñ$ˆ
�A�v˜uØ—_‘_ Z°×1BÑ1BÀAÀvÈuÓUˆ
ð '¢qª!¨Q° |Ñ4Ðä—m‘m×/Ñ/°¸6À5¸/ÐPZÐjnÐ/Óoˆð ×$Ò$Ü!ŸJ™J¬u¯y©yÑ9¸$¿)¹)ÑDˆDáœ=Ð):×)DÑ)DÀQÓ)GÈ+×J_ÑJ_Ð`aÓJbÑ)bÓcÐijÔk‰Gä×&Ñ& {¸;×;MÑ;MÔNˆGÜ˜4×,Ñ,Ó-ò b�àœ=Ð):º1¸aÀ¸9Ñ)E×)OÑ)OÐPQÓ)RÐU`Ñ)`ÓaÑa‘ðbð �y‰y˜FÒ"Ø! D×$5Ñ$5Ñ5�à%¨Ñ/ˆØ—~‘~¨¯©Ñ6¸/ÑIÈDÏNÉNÑZˆÜŸ™ K°QÔ7‰ˆ�QÜ—j‘j ¨d¯n©n¸d¿n¹nÓMˆØ Ð+Ð+r$   )r  r  r  F©NTrs   ra   s   @r%   r  r  Q  s   ø„ ð
 ØØØõ*÷><,r$   r  c                   ó0   ‡ — e Zd Z	 	 	 	 dˆ fd„	Zdd„Zˆ xZS )ÚZoeDepthAttractorLayerUnnormedc                 óÊ  •— t         ‰	| �  «        || _        || _        || _        || _        |j                  | _        |j                  | _        |j                  | _
        || _        |j                  x}}t        j                  ||ddd«      | _        t        j                   d¬«      | _        t        j                  ||ddd«      | _        t        j&                  «       | _        y)zL
        Attractor layer for bin centers. Bin centers are unbounded
        r   r   Trö   N)r-   r.   r   rÿ   rû   rü   r  r  r  r  r  r!  r"  r   rj   r¼   r†   rý   rÀ   rä   rþ   r#  s
            €r%   r.   z'ZoeDepthAttractorLayerUnnormed.__init__±  s¹   ø€ ô 	‰ÑÔà(ˆÔØˆŒØ"ˆŒØ"ˆŒØ×+Ñ+ˆŒ
Ø×+Ñ+ˆŒ
Ø×)Ñ)ˆŒ	Ø 0ˆÔà &× 8Ñ 8Ð8ˆ�gÜ—Y‘Y˜{¨G°Q¸¸1Ó=ˆŒ
Ü—G‘G DÔ)ˆŒ	Ü—Y‘Y˜w¨°a¸¸AÓ>ˆŒ
Ü—K‘K“Mˆ�	r$   c                 ó`  — |�7|r0t         j                  j                  ||j                  dd dd¬«      }||z   }| j	                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }|j                  dd \  }}t         j                  j                  |||fdd¬«      }| j                  sct        j                  t        j                  dœ| j                     }	 |	t        |j                  d«      |j                  d«      z
  «      d¬	«      }
n�t        j                  ||j                   ¬
«      }
t#        | j$                  «      D ]*  }|
t        |dd…|df   j                  d«      |z
  «      z  }
Œ, | j                  dk(  r|
| j$                  z  }
||
z   }|}||fS )a¤  
        The forward pass of the attractor layer. This layer predicts the new bin centers based on the previous bin centers
        and the attractor points (the latter are predicted by the MLP).

        Args:
            x (`torch.Tensor` of shape (batch_size, num_channels, height, width)`):
                Feature block.
            prev_bin (`torch.Tensor` of shape (batch_size, prev_num_bins, height, width)`):
                Previous bin centers normed.
            prev_bin_embedding (`torch.Tensor`, *optional*):
                Optional previous bin embeddings.
            interpolate (`bool`, *optional*, defaults to `True`):
                Whether to interpolate the previous bin embeddings to the size of the input features.

        Returns:
            `Tuple[`torch.Tensor`, `torch.Tensor`]:
                New bin centers unbounded. Two outputs just to keep the API consistent with the normed version.
        Nr%  r™   Tr&  r'  r,   r   rD   r)  .r(  )r   rŸ   r    rH   r¼   rý   rÀ   rþ   r!  r    r(  r  r  r  rN   r+  r*  rw   r   )r<   r
  r.  r/  r    r0  r1  r2  r  r4  r5  r²   r6  s                r%   r]   z&ZoeDepthAttractorLayerUnnormed.forwardÎ  s�  € ð& Ð)ÙÜ%'§]¡]×%>Ñ%>Ø&¨¯©°°¨¸:ÐUYð &?ó &Ð"ð Ð&Ñ&ˆAà�J‰J�q‹MˆØ�I‰I�a‹LˆØ�J‰J�q‹MˆØ—Y‘Y˜q“\ˆ
à"×(Ñ(¨¨Ð-‰ˆ�ä—m‘m×/Ñ/°¸6À5¸/ÐPZÐjnÐ/Óoˆà×$Ò$Ü!ŸJ™J¬u¯y©yÑ9¸$¿)¹)ÑDˆDáœ=¨×)=Ñ)=¸aÓ)@À;×CXÑCXÐYZÓC[Ñ)[Ó\ÐbcÔd‰Gä×&Ñ& {¸;×;MÑ;MÔNˆGÜ˜4×,Ñ,Ó-ò [�àœ=¨²A°q¸#°IÑ)>×)HÑ)HÈÓ)KÈkÑ)YÓZÑZ‘ð[ð �y‰y˜FÒ"Ø! D×$5Ñ$5Ñ5�à%¨Ñ/ˆØ%ˆà Ð+Ð+r$   )r  r  r  Tr7  rs   ra   s   @r%   r9  r9  °  s   ø„ ð
 ØØØõ"÷:3,r$   r9  c                   óX   ‡ — e Zd Zdˆ fd„	Zdej
                  dej
                  fd„Zˆ xZS )ÚZoeDepthProjectorc                 óÐ   •— t         ‰| �  «        t        j                  ||ddd«      | _        t        j
                  d¬«      | _        t        j                  ||ddd«      | _        y)a  Projector MLP.

        Args:
            in_features (`int`):
                Number of input channels.
            out_features (`int`):
                Number of output channels.
            mlp_dim (`int`, *optional*, defaults to 128):
                Hidden dimension.
        r   r   Trö   N)r-   r.   r   rj   r¼   r†   rÏ   rÀ   )r<   rê   Úout_featuresr   rA   s       €r%   r.   zZoeDepthProjector.__init__  sP   ø€ ô 	‰ÑÔä—Y‘Y˜{¨G°Q¸¸1Ó=ˆŒ
Ü—7‘7 4Ô(ˆŒÜ—Y‘Y˜w¨°a¸¸AÓ>ˆ�
r$   r\   rB   c                 ól   — | j                  |«      }| j                  |«      }| j                  |«      }|S rq   )r¼   rÏ   rÀ   rr   s     r%   r]   zZoeDepthProjector.forward  s2   € Ø—z‘z ,Ó/ˆØ—x‘x Ó-ˆØ—z‘z ,Ó/ˆàÐr$   )é€   )r   r   r   r.   r    r_   r]   r`   ra   s   @r%   r=  r=    s#   ø„ õ?ð" E§L¡Lð °U·\±\÷ r$   r=  c                   óö   ‡ — e Zd ZdZˆ fd„Zdej                  dej                  fd„Z	 	 ddej                  dej                  dej                  d	eej                     d
ee
   deej                     fd„Zˆ xZS )ÚZoeDepthMultiheadAttentionzKEquivalent implementation of nn.MultiheadAttention with `batch_first=True`.c                 ó  •— t         ‰| �  «        ||z  dk7  rt        d|› d|› d�«      ‚|| _        t	        ||z  «      | _        | j                  | j
                  z  | _        t        j                  || j                  «      | _	        t        j                  || j                  «      | _
        t        j                  || j                  «      | _        t        j                  ||«      | _        t        j                  |«      | _        y )Nr   zThe hidden size (z6) is not a multiple of the number of attention heads (ú))r-   r.   r±   Únum_attention_headsro   Úattention_head_sizeÚall_head_sizer   r:   ÚqueryÚkeyr  Úout_projÚDropoutÚdropout)r<   r?   rF  rM  rA   s       €r%   r.   z#ZoeDepthMultiheadAttention.__init__#  sâ   ø€ Ü‰ÑÔØÐ,Ñ,°Ò1ÜØ# K =ð 1Ø-Ð.¨að1óð ð
 $7ˆÔ Ü#& {Ð5HÑ'HÓ#IˆÔ Ø!×5Ñ5¸×8PÑ8PÑPˆÔä—Y‘Y˜{¨D×,>Ñ,>Ó?ˆŒ
Ü—9‘9˜[¨$×*<Ñ*<Ó=ˆŒÜ—Y‘Y˜{¨D×,>Ñ,>Ó?ˆŒ
äŸ	™	 +¨{Ó;ˆŒä—z‘z 'Ó*ˆ�r$   r
  rB   c                 ó¤   — |j                  «       d d | j                  | j                  fz   }|j                  |«      }|j	                  dddd«      S )NrF   r   r,   r   r	   )r›   rF  rG  rÒ   rK   )r<   r
  Únew_x_shapes      r%   Útranspose_for_scoresz/ZoeDepthMultiheadAttention.transpose_for_scores7  sL   € Ø—f‘f“h˜s �m t×'?Ñ'?À×AYÑAYÐ&ZÑZˆØ�F‰F�;ÓˆØ�y‰y˜˜A˜q !Ó$Ð$r$   ÚqueriesÚkeysÚvaluesÚattention_maskÚoutput_attentionsc                 óÔ  — | j                  | j                  |«      «      }| j                  | j                  |«      «      }| j                  | j                  |«      «      }t	        j
                  ||j                  dd«      «      }	|	t        j                  | j                  «      z  }	|�|	|z   }	t        j                  j                  |	d¬«      }
| j                  |
«      }
t	        j
                  |
|«      }|j                  dddd«      j                  «       }|j!                  «       d d | j"                  fz   }|j%                  |«      }| j'                  |«      }|r||
f}|S |f}|S )NrF   r%  rD   r   r,   r   r	   )rP  rI  rJ  r  r    ÚmatmulÚ	transposeÚmathÚsqrtrG  r   rŸ   rÞ   rM  rK   rL   r›   rH  rÒ   rK  )r<   rQ  rR  rS  rT  rU  Úquery_layerÚ	key_layerÚvalue_layerÚattention_scoresÚattention_probsÚcontext_layerÚnew_context_layer_shapeÚoutputss                 r%   r]   z"ZoeDepthMultiheadAttention.forward<  sY  € ð ×/Ñ/°·
±
¸7Ó0CÓDˆØ×-Ñ-¨d¯h©h°t«nÓ=ˆ	Ø×/Ñ/°·
±
¸6Ó0BÓCˆô !Ÿ<™<¨°Y×5HÑ5HÈÈRÓ5PÓQÐà+¬d¯i©i¸×8PÑ8PÓ.QÑQÐØÐ%à/°.Ñ@Ðô Ÿ-™-×/Ñ/Ð0@ÀbÐ/ÓIˆð Ÿ,™, Ó7ˆäŸ™ _°kÓBˆà%×-Ñ-¨a°°A°qÓ9×DÑDÓFˆØ"/×"4Ñ"4Ó"6°s¸Ð";¸t×?QÑ?QÐ>SÑ"SÐØ%×*Ñ*Ð+BÓCˆàŸ™ mÓ4ˆá6G�= /Ð2ˆàˆð O\ÐM]ˆàˆr$   )NF)r   r   r   r   r.   r    r_   rP  r   r!   Úboolr   r]   r`   ra   s   @r%   rC  rC    s‘   ø„ ÙUô+ð(% e§l¡lð %°u·|±|ó %ð 7;Ø,1ñ%à—‘ð%ð �l‰lð%ð —‘ð	%ð
 ! ×!2Ñ!2Ñ3ð%ð $ D™>ð%ð 
ˆu�|‰|Ñ	÷%r$   rC  c                   óJ   ‡ — e Zd Zdˆ fd„	Z	 ddeej                     fd„Zˆ xZS )ÚZoeDepthTransformerEncoderLayerc                 ó  •— t         ‰| �  «        |j                  }|j                  }|j                  }t        |||¬«      | _        t        j                  ||«      | _	        t        j                  |«      | _        t        j                  ||«      | _        t        j                  |«      | _        t        j                  |«      | _        t        j                  |«      | _        t        j                  |«      | _        t$        |   | _        y )N)rM  )r-   r.   Úpatch_transformer_hidden_sizeÚ#patch_transformer_intermediate_sizeÚ%patch_transformer_num_attention_headsrC  Ú	self_attnr   r:   Úlinear1rL  rM  Úlinear2Ú	LayerNormÚnorm1Únorm2Údropout1Údropout2r
   Ú
activation)r<   r=   rM  rr  r?   Úintermediate_sizerF  rA   s          €r%   r.   z(ZoeDepthTransformerEncoderLayer.__init__e  sÅ   ø€ Ü‰ÑÔà×:Ñ:ˆØ"×FÑFÐØ$×JÑJÐä3°KÐATÐ^eÔfˆŒä—y‘y Ð.?Ó@ˆŒÜ—z‘z 'Ó*ˆŒÜ—y‘yÐ!2°KÓ@ˆŒä—\‘\ +Ó.ˆŒ
Ü—\‘\ +Ó.ˆŒ
ÜŸ
™
 7Ó+ˆŒÜŸ
™
 7Ó+ˆŒä  Ñ,ˆ�r$   Úsrc_maskc           	      óN  — |x}}| j                  ||||¬«      d   }|| j                  |«      z   }| j                  |«      }| j                  | j	                  | j                  | j                  |«      «      «      «      }|| j                  |«      z   }| j                  |«      }|S )N)rQ  rR  rS  rT  r   )	rj  rp  rn  rl  rM  rr  rk  rq  ro  )r<   Úsrcrt  rQ  rR  Úsrc2s         r%   r]   z'ZoeDepthTransformerEncoderLayer.forwardy  s™   € ð
 Ðˆ�$Ø�~‰~ g°DÀÐU]ˆ~Ó^Ð_`ÑaˆØ�D—M‘M $Ó'Ñ'ˆØ�j‰j˜‹oˆØ�|‰|˜DŸL™L¨¯©¸¿¹ÀcÓ9JÓ)KÓLÓMˆØ�D—M‘M $Ó'Ñ'ˆØ�j‰j˜‹oˆØˆ
r$   )gš™™™™™¹?Úrelurq   )	r   r   r   r.   r   r    r_   r]   r`   ra   s   @r%   re  re  d  s%   ø„ õ-ð. ,0ñð ˜5Ÿ<™<Ñ(÷r$   re  c                   óD   ‡ — e Zd Zˆ fd„Zdej
                  fd„Zd„ Zˆ xZS )ÚZoeDepthPatchTransformerEncoderc                 ó  •— t         ‰| �  «        |j                  }t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        t        j                  ||j                  ddd¬«      | _        yc c}w )z¤ViT-like transformer block

        Args:
            config (`ZoeDepthConfig`):
                Model configuration class defining the model architecture.
        r   r   rg   N)r-   r.   rù   r   r0   rw   Únum_patch_transformer_layersre  Útransformer_encoderrj   rg  Úembedding_convPxP)r<   r=   rd   r@   rA   s       €r%   r.   z(ZoeDepthPatchTransformerEncoder.__init__‰  sw   ø€ ô 	‰ÑÔà×0Ñ0ˆä#%§=¡=Ü>CÀF×DgÑDgÓ>hÖi¸Ô,¨VÕ4Òió$
ˆÔ ô "$§¡Ø˜×=Ñ=È1ÐUVÐ`aô"
ˆÕùò js   ÁB
Úcpuc           	      óþ  — t        j                  d|||¬«      j                  d«      }t        j                  d|d||¬«      j                  d«      }t        j                  |t        j                  t        j
                  d|¬«      «       |z  z  «      }||z  }	t        j                  t        j                  |	«      t        j                  |	«      gd¬«      }	|	j                  d¬«      j                  |dd«      }	|	S )zßGenerate positional encodings

        Args:
            sequence_length (int): Sequence length
            embedding_dim (int): Embedding dimension

        Returns:
            torch.Tensor: Positional encodings.
        r   )Údtyper*  r   r,   g     ˆÃ@r)  rD   )
r    rÑ   rN   ÚexprÅ   rÓ   rI   ÚsinÚcosÚrepeat)
r<   rT   rW   Úembedding_dimr*  r�  ÚpositionÚindexÚdiv_termÚpos_encodings
             r%   Úpositional_encoding_1dz6ZoeDepthPatchTransformerEncoder.positional_encoding_1dœ  sÐ   € ô —<‘<  ?¸%ÈÔO×YÑYÐZ[Ó\ˆÜ—‘˜Q ¨q¸ÀfÔM×WÑWÐXYÓZˆÜ—9‘9˜U¤u§y¡y´·±¸gÈfÔ1UÓ'VÐ&VÐYfÑ&fÑgÓhˆØ (Ñ*ˆÜ—y‘y¤%§)¡)¨LÓ"9¼5¿9¹9À\Ó;RÐ!SÐYZÔ[ˆØ#×-Ñ-°!Ð-Ó4×;Ñ;¸JÈÈ1ÓMˆØÐr$   c                 óp  — | j                  |«      j                  d«      }t        j                  j	                  |d«      }|j                  ddd«      }|j                  \  }}}|| j                  ||||j                  |j                  ¬«      z   }t        d«      D ]  } | j                  |   |«      }Œ |S )zßForward pass

        Args:
            x (torch.Tensor - NCHW): Input feature tensor

        Returns:
            torch.Tensor - Transformer output embeddings of shape (batch_size, sequence_length, embedding_dim)
        r,   )r   r   r   r   )r*  r�  râ   )r~  rM   r   rŸ   r  rK   rH   r‹  r*  r�  rw   r}  )r<   r
  Ú
embeddingsrT   rW   r†  r²   s          r%   r]   z'ZoeDepthPatchTransformerEncoder.forward®  sÅ   € ð ×+Ñ+¨AÓ.×6Ñ6°qÓ9ˆ
ä—]‘]×&Ñ& z°6Ó:ˆ
à×'Ñ'¨¨1¨aÓ0ˆ
Ø5?×5EÑ5EÑ2ˆ
�O ]Ø $×"=Ñ"=Ø˜¨¸z×?PÑ?PÐXb×XhÑXhð #>ó #
ñ 
ˆ
ô �q“ò 	AˆAØ4˜×1Ñ1°!Ñ4°ZÓ@‰Jð	Að Ðr$   )	r   r   r   r.   r    Úfloat32r‹  r]   r`   ra   s   @r%   rz  rz  ˆ  s"   ø„ ô
ð& Y^Ðej×erÑeró ö$r$   rz  c                   ó&   ‡ — e Zd Zdˆ fd„Zd„ Zˆ xZS )ÚZoeDepthMLPClassifierc                 óÄ   •— t         ‰| �  «        |}t        j                  ||«      | _        t        j
                  «       | _        t        j                  ||«      | _        y rq   )r-   r.   r   r:   rk  r†   rr  rl  )r<   rê   r?  Úhidden_featuresrA   s       €r%   r.   zZoeDepthMLPClassifier.__init__È  sD   ø€ Ü‰ÑÔà%ˆÜ—y‘y ¨oÓ>ˆŒÜŸ'™'›)ˆŒÜ—y‘y °,Ó?ˆ�r$   c                 ól   — | j                  |«      }| j                  |«      }| j                  |«      }|S rq   )rk  rr  rl  )r<   r\   r   s      r%   r]   zZoeDepthMLPClassifier.forwardÐ  s2   € Ø—|‘| LÓ1ˆØ—‘ |Ó4ˆØŸ™ \Ó2ˆàÐr$   )rB   Nrs   ra   s   @r%   r�  r�  Ç  s   ø„ õ@ör$   r�  c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )Ú*ZoeDepthMultipleMetricDepthEstimationHeadszn
    Multiple metric depth estimation heads. A MLP classifier is used to route between 2 different heads.
    c                 ó¼  •— t         ‰| �  «        |j                  }|j                  }|j                  | _        |j
                  | _        |j                  }t        j                  ||ddd¬«      | _	        t        |«      | _        t        dd¬«      | _        | j
                  dk(  rt        }n| j
                  dk(  rt        }t        j                   |j                  D �ci c]"  }|d	   t#        ||d
   |dz  |d   |d   ¬«      “Œ$ c}«      | _        t'        |||dz  ¬«      | _        t        j*                  t-        d«      D �cg c]  }t'        |j.                  ||dz  ¬«      ‘Œ c}«      | _        t        j                   |j                  D ��	ci c]N  }|d	   t        j*                  t-        t3        |«      «      D �	cg c]  }	 |||	   |d   |d   ¬«      ‘Œ c}	«      “ŒP c}	}«      | _        |j6                  }
t        j                   |j                  D �ci c]  }|d	   t9        ||
||d
   d¬«      “Œ c}«      | _        y c c}w c c}w c c}	w c c}	}w c c}w )Nr   r   rg   rA  r,   ©rê   r?  rø   ÚsoftplusÚnamerÿ   rû   rü   )rÿ   r   rû   rü   )rê   r?  r   râ   ©rÿ   rû   rü   )rì   )r-   r.   r"  Únum_attractorsÚbin_configurationsrú   rù   r   rj   rÀ   rz  Úpatch_transformerr�  Úmlp_classifierr  r9  Ú
ModuleDictrô   Úseed_bin_regressorsr=  Úseed_projectorr0   rw   rˆ   Ú
projectorsrx   r0  r¿   rà   Úconditional_log_binomial)r<   r=   r"  r   rù   Ú	AttractorÚconfr@   Úconfigurationr²   Úlast_inrA   s              €r%   r.   z3ZoeDepthMultipleMetricDepthEstimationHeads.__init__Ý  s‰  ø€ Ü‰ÑÔà"×4Ñ4ÐØ×,Ñ,ˆØ"(×";Ñ";ˆÔØ &× 7Ñ 7ˆÔð %×8Ñ8ÐÜ—Y‘YÐ2Ð4GÐUVÐ_`ÐjkÔlˆŒ
ô "AÀÓ!HˆÔä3ÀÐRSÔTˆÔð × Ñ  HÒ,Ü.‰IØ×"Ñ" jÒ0Ü6ˆIô $&§=¡=ð #×5Ñ5ö	ð ð �V‘Ô6ØØ ™>Ø-°Ñ2Ø" ;Ñ/Ø" ;Ñ/ôñ ò	ó$
ˆÔ ô 0Ø+Ð:KÐUfÐjkÑUkô
ˆÔô Ÿ-™-ô ˜q›öð ô "Ø &× 9Ñ 9Ø!2Ø-°Ñ2öòó	
ˆŒô Ÿ-™-ð &,×%>Ñ%>÷ð "ð ˜fÑ%¤r§}¡}ô "'¤s¨<Ó'8Ó!9öð ñ "Ø"Ø#/°¡?Ø&3°KÑ&@Ø&3°KÑ&@ö	òó
(ñ 
óó
ˆŒð" ×.Ñ.ˆä(*¯©ð &,×%>Ñ%>ö	ð "ð ˜fÑ%Ô'LØØØ%Ø! (Ñ+Ø&'ô(ñ ò	ó)
ˆÕ%ùò]	ùò ùòùóùò&	s*   Ã'IÅ"I	Æ.I
Æ>IÇ	I
ÈIÉI
c                 ór  — | j                  |«      }| j                  |«      d d …dd d …f   }| j                  |«      }t        j                  |j                  dd¬«      d¬«      }| j                  D �	cg c]  }	|	d   ‘Œ	 }
}	|
t        j                  |d¬«      j                  «       j                  «          }	 | j                  D �cg c]  }|d   |k(  sŒ|‘Œ c}d   }|d	   }|d
   }| j                  |   } ||«      \  }}| j                  dv r||z
  ||z
  z  }n|}| j                  |«      }| j                  |   }t!        | j"                  ||«      D ]!  \  }}} ||«      } ||||d¬«      \  }}|}|}Œ# |}t$        j&                  j)                  |j*                  dd  dd¬«      }t$        j&                  j)                  |j*                  dd  dd¬«      }| j,                  |   } |||«      }t        j
                  ||z  dd¬«      }||fS c c}	w c c}w # t        $ r t        d|› d�«      ‚w xY w)Nr   Tr  rF   rD   r™  zbin_configurations_name z! not found in bin_configurationssrû   rü   ©rø   Úhybrid2©r    r%  r™   r&  r   )rÀ   r�  rž  r    rÞ   r  rœ  ÚargmaxrÃ   ÚitemÚ
IndexErrorr±   r   rú   r¡  r0  r2   r¢  r   rŸ   r    rH   r£  )r<   Úoutconv_activationrí   Úfeature_blocksÚrelative_depthr
  Ú	embeddingr   Údomain_voter¦  ÚnamesÚbin_configurations_namer=   r¥  rû   rü   Úseed_bin_regressorr@   Úseed_bin_centersr.  r/  r0  Ú	projectorÚ	attractorr³   Úbin_embeddingÚbinr  Úlastr£  rZ   s                                  r%   r]   z2ZoeDepthMultipleMetricDepthEstimationHeads.forward1  sv  € Ø�J‰J�zÓ"ˆð ×*Ñ*¨1Ó-ªa°²A¨gÑ6ˆ	ð ×+Ñ+¨IÓ6ˆÜ—m‘m M×$5Ñ$5¸!ÀTÐ$5Ó$JÐPRÔSˆð =A×<SÑ<SÖT¨=�˜vÓ&ÐTˆÐTØ"'¬¯©°[ÀbÔ(I×(QÑ(QÓ(S×(XÑ(XÓ(ZÑ"[Ðð	tØ)-×)@Ñ)@Ön˜vÀFÈ6ÁNÐVmÓDm’FÒnÐopÑqˆDð ˜Ñ%ˆ	Ø˜Ñ%ˆ	à!×5Ñ5Ð6MÑNÐÙ0°Ó3ÑˆÐØ× Ñ Ð$9Ñ9Ø(¨9Ñ4¸ÀYÑ9NÑO‰Hà'ˆHØ!×0Ñ0°Ó3Ðà—_‘_Ð%<Ñ=ˆ
Ü-0°·±À*ÈnÓ-]ò 	/Ñ)ˆI�y 'Ù% gÓ.ˆMÙ(¨¸ÐBTÐbfÔgÑˆC�ØˆHØ!.Ñð		/ð "ˆä—m‘m×/Ñ/°¸T¿Z¹ZÈÈ¸_ÐS]ÐmqÐ/ÓrˆÜŸ™×1Ñ1°-ÀÇÁÈBÈCÀÐWaÐquÐ1Óvˆà#'×#@Ñ#@ÐAXÑ#YÐ Ù$ T¨=Ó9ˆô �i‰i˜˜K™¨Q¸Ô=ˆà�MÐ!Ð!ùòK Uùò oøÜò 	tÜÐ7Ð8OÐ7PÐPqÐrÓsÐsð	tús*   Á4HÂ9H ÃHÃHÃH ÈH ÈH6r¡   ra   s   @r%   r•  r•  Ø  s   ø„ ñôR
öh1"r$   r•  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )Ú!ZoeDepthMetricDepthEstimationHeadc                 ó,  •— t         ‰| �  «        |j                  d   }|d   }|d   }|d   }|j                  }|j                  }|j
                  }|| _        || _        || _        |j                  }	t        j                  |	|	ddd¬«      | _        | j
                  dk(  rt        }
n| j
                  dk(  rt        }
t        ||||¬	«      | _        t!        |	|¬
«      | _        t        j$                  t'        d«      D �cg c]  }t!        |j(                  |¬
«      ‘Œ c}«      | _        t        j$                  t'        d«      D �cg c]  } 
||||   ||¬«      ‘Œ c}«      | _        |j.                  dz   }t1        ||||¬«      | _        y c c}w c c}w )Nr   rÿ   rû   rü   r   rg   rø   r˜  rš  r—  râ   )rÿ   r   rû   rü   )rÔ   )r-   r.   rœ  r"  r›  rú   rû   rü   rù   r   rj   rÀ   r  r9  rô   r¶  r=  r¡  r0   rw   rˆ   r¢  r0  r¿   rà   r£  )r<   r=   Úbin_configurationrÿ   rû   rü   r"  r   rú   rù   r¤  r@   r²   r§  rA   s                 €r%   r.   z*ZoeDepthMetricDepthEstimationHead.__init__f  s¬  ø€ Ü‰ÑÔà"×5Ñ5°aÑ8ÐØ" 8Ñ,ˆØ% kÑ2ˆ	Ø% kÑ2ˆ	Ø"×4Ñ4ÐØ×,Ñ,ˆØ!×2Ñ2Ðà"ˆŒØ"ˆŒØ 0ˆÔð %×8Ñ8ÐÜ—Y‘YÐ2Ð4GÐUVÐ_`ÐjkÔlˆŒ
ð × Ñ  HÒ,Ü.‰IØ×"Ñ" jÒ0Ü6ˆIä":Ø˜6¨YÀ)ô#
ˆÔô 0Ð<OÐ^oÔpˆÔäŸ-™-ô ˜q›öàô "¨f×.GÑ.GÐVgÖhòó
ˆŒô Ÿ-™-ô ˜q›ö	ð ñ ØØ!Ø!-¨a¡Ø'Ø'öò	ó
ˆŒð ×.Ñ.°Ñ2ˆô )NØØØØô	)
ˆÕ%ùò+ùò	s   Ã?FÅFc                 ó~  — | j                  |«      }| j                  |«      \  }}| j                  dv r*|| j                  z
  | j                  | j                  z
  z  }n|}| j                  |«      }	t        | j                  | j                  |«      D ]=  \  }
}} |
|«      } ||||	d¬«      \  }}|j                  «       }|j                  «       }	Œ? |}|j                  d«      }t        j                  j                  ||j                  dd  dd¬«      }t        j                   ||gd¬«      }t        j                  j                  |j                  d	d  dd¬
«      }| j#                  ||«      }t        j                  j                  |j                  d	d  dd¬
«      }t        j$                  ||z  dd¬«      }|d fS )Nr©  Tr«  r   r,   r™   rš   rD   r%  r&  r  )rÀ   r¶  rú   rû   rü   r¡  r2   r¢  r0  ÚclonerN   r   rŸ   r    rH   r    rI   r£  r  )r<   r¯  rí   r°  r±  r
  r@   r·  r.  r/  r¸  r¹  r³   rº  r»  r  r¼  Úrelative_conditioningrZ   s                      r%   r]   z)ZoeDepthMetricDepthEstimationHead.forward¡  s¶  € Ø�J‰J�zÓ"ˆØ"×5Ñ5°aÓ8ÑˆÐà× Ñ Ð$9Ñ9Ø(¨4¯>©>Ñ9¸d¿n¹nÈtÏ~É~Ñ>]Ñ^‰Hà'ˆHà!×0Ñ0°Ó3Ðô .1°·±À$Ç/Á/ÐSaÓ-bò 	7Ñ)ˆI�y 'Ù% gÓ.ˆMÙ(¨¸ÐBTÐbfÔgÑˆC�Ø—y‘y“{ˆHØ!.×!4Ñ!4Ó!6Ñð		7ð "ˆð !/× 8Ñ 8¸Ó ;ÐÜ "§¡× 9Ñ 9Ø!¨¯
©
°1°2¨¸ZÐW[ð !:ó !
Ðô �y‰y˜$Ð 5Ð6¸AÔ>ˆäŸ™×1Ñ1°-ÀÇÁÈBÈCÀÐWaÐquÐ1ÓvˆØ×)Ñ)¨$°Ó>ˆô —m‘m×/Ñ/°¸Q¿W¹WÀRÀS¸\ÐPZÐjnÐ/ÓoˆÜ�i‰i˜˜K™¨Q¸Ô=ˆà�DˆyÐr$   rs   ra   s   @r%   r¾  r¾  e  s   ø„ ô9
öv"r$   r¾  c                   ó&   — e Zd ZdZeZdZdZdZd„ Z	y)ÚZoeDepthPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    ÚzoedepthÚpixel_valuesTc                 ó  — t        |t        j                  t        j                  t        j                  f«      rm|j
                  j                  j                  d| j                  j                  ¬«       |j                  �%|j                  j                  j                  «        yyt        |t        j                  «      rJ|j                  j                  j                  «        |j
                  j                  j                  d«       yy)zInitialize the weightsg        )r(  ÚstdNrÜ   )r­   r   r:   rj   rl   ÚweightÚdataÚnormal_r=   Úinitializer_ranger�   Úzero_rm  Úfill_)r<   Úmodules     r%   Ú_init_weightsz%ZoeDepthPreTrainedModel._init_weightsÓ  s°   € ä�fœrŸy™y¬"¯)©)´R×5GÑ5GÐHÔIð �M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ä˜¤§¡Ô-Ø�K‰K×Ñ×"Ñ"Ô$Ø�M‰M×Ñ×$Ñ$ SÕ)ð .r$   N)
r   r   r   r   r   Úconfig_classÚbase_model_prefixÚmain_input_nameÚsupports_gradient_checkpointingrÑ  r#   r$   r%   rÅ  rÅ  È  s$   „ ñð
 "€LØ"ÐØ$€OØ&*Ð#ó
*r$   rÅ  aE  
    This model is a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass. Use it
    as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage and
    behavior.

    Parameters:
        config ([`ViTConfig`]): Model configuration class with all the parameters of the model.
            Initializing with a config file does not load the weights associated with the model, only the
            configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.
a  
    Args:
        pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See [`DPTImageProcessor.__call__`]
            for details.

        output_attentions (`bool`, *optional*):
            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
            tensors for more detail.
        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~file_utils.ModelOutput`] instead of a plain tuple.
zU
    ZoeDepth model with one or multiple metric depth estimation head(s) on top.
    c                   óÜ   ‡ — e Zd Zˆ fd„Z ee«       eee¬«      	 	 	 	 d
de	j                  dee	j                     dee   dee   dee   deee	j                      ef   fd	„«       «       Zˆ xZS )ÚZoeDepthForDepthEstimationc                 ó6  •— t         ‰| �  |«       t        |«      | _        t	        | j                  j
                  d«      rkt	        | j                  j
                  d«      rK| j                  j
                  j                  |_        | j                  j
                  j                  | _        nt        d«      ‚t        |«      | _        t        |«      | _        t        |j                  «      dkD  rt!        |«      n
t#        |«      | _        | j'                  «        y )Nr?   Ú
patch_sizezXZoeDepth assumes the backbone's config to have `hidden_size` and `patch_size` attributesr   )r-   r.   r   ÚbackboneÚhasattrr=   r?   r8   rÙ  r±   r£   Úneckr·   Úrelative_headrx   rœ  r•  r¾  Úmetric_headÚ	post_init)r<   r=   rA   s     €r%   r.   z#ZoeDepthForDepthEstimation.__init__  sÙ   ø€ Ü‰Ñ˜Ô ä% fÓ-ˆŒä�4—=‘=×'Ñ'¨Ô7¼GÀDÇMÁM×DXÑDXÐZfÔ<gØ*.¯-©-×*>Ñ*>×*JÑ*JˆFÔ'Ø"Ÿm™m×2Ñ2×=Ñ=ˆD�OäØjóð ô ! Ó(ˆŒ	Ü@ÀÓHˆÔô �6×,Ñ,Ó-°Ò1ô 7°vÔ>ä2°6Ó:ð 	Ôð 	�‰Õr$   )Úoutput_typerÒ  rÇ  ÚlabelsrU  Úoutput_hidden_statesÚreturn_dictrB   c                 ó¼  — d}|�t        d«      ‚|�|n| j                  j                  }|�|n| j                  j                  }|�|n| j                  j                  }| j
                  j                  |||¬«      }|j                  }|j                  \  }	}	}
}| j                  }|
|z  }||z  }| j                  |||«      \  }}|g|z   }| j                  |«      \  }}|g|z   }| j                  |d   |d   |dd |¬«      \  }}|j                  d¬«      }|s |�||f|dd z   }n	|f|dd z   }|�|f|z   S |S t        ||||j                  |j                   ¬	«      S )
a¶  
        labels (`torch.LongTensor` of shape `(batch_size, height, width)`, *optional*):
            Ground truth depth estimation maps for computing the loss.

        Returns:

        Examples:
        ```python
        >>> from transformers import AutoImageProcessor, ZoeDepthForDepthEstimation
        >>> import torch
        >>> import numpy as np
        >>> from PIL import Image
        >>> import requests

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = Image.open(requests.get(url, stream=True).raw)

        >>> image_processor = AutoImageProcessor.from_pretrained("Intel/zoedepth-nyu-kitti")
        >>> model = ZoeDepthForDepthEstimation.from_pretrained("Intel/zoedepth-nyu-kitti")

        >>> # prepare image for the model
        >>> inputs = image_processor(images=image, return_tensors="pt")

        >>> with torch.no_grad():
        ...     outputs = model(**inputs)

        >>> # interpolate to original size
        >>> post_processed_output = image_processor.post_process_depth_estimation(
        ...     outputs,
        ...     source_sizes=[(image.height, image.width)],
        ... )

        >>> # visualize the prediction
        >>> predicted_depth = post_processed_output[0]["predicted_depth"]
        >>> depth = predicted_depth * 255 / predicted_depth.max()
        >>> depth = depth.detach().cpu().numpy()
        >>> depth = Image.fromarray(depth.astype("uint8"))
        ```NzTraining is not implemented yet)râ  rU  r   r   r,   )r¯  rí   r°  r±  rD   )r   r   r   r   r   )ÚNotImplementedErrorr=   Úuse_return_dictrâ  rU  rÚ  Úforward_with_filtered_kwargsÚfeature_mapsrH   rÙ  rÜ  rÝ  rÞ  rÃ   r   r   r   )r<   rÇ  rá  rU  râ  rã  r   rb  r   r@   r1  r2  rÙ  rR   rS   r´   rZ   r±  Úmetric_depthr   rµ   s                        r%   r]   z"ZoeDepthForDepthEstimation.forward  sÄ  € ð` ˆØÐÜ%Ð&GÓHÐHà%0Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà—-‘-×<Ñ<ØÐ/CÐWhð =ó 
ˆð  ×,Ñ,ˆà*×0Ñ0Ñˆˆ1ˆf�eØ—_‘_ˆ
Ø Ñ+ˆØ˜zÑ)ˆà"&§)¡)¨M¸<ÈÓ"UÑˆ�xàˆj˜=Ñ(ˆà#'×#5Ñ#5°mÓ#DÑ ˆ˜àˆj˜3Ñˆà&*×&6Ñ&6Ø" 1™v°#°a±&ÈÈQÈRÈÐaoð '7ó '
Ñ#ˆ�mð $×+Ñ+°Ð+Ó2ˆáØÐ(Ø&¨Ð6¸ÀÀ¸ÑD‘à&˜¨7°1°2¨;Ñ6�à)-Ð)9�T�G˜fÑ$ÐE¸vÐEä+ØØ(Ø'Ø!×/Ñ/Ø×)Ñ)ô
ð 	
r$   )NNNN)r   r   r   r.   r   ÚZOEDEPTH_INPUTS_DOCSTRINGr   r   Ú_CONFIG_FOR_DOCr    r!   r   Ú
LongTensorrc  r   r   r_   r]   r`   ra   s   @r%   r×  r×  ü  sµ   ø„ ôñ2 +Ð+DÓEÙÐ+?ÈoÔ^ð .2Ø,0Ø/3Ø&*ñ]
à×'Ñ'ð]
ð ˜×)Ñ)Ñ*ð]
ð $ D™>ð	]
ð
 ' t™nð]
ð ˜d‘^ð]
ð 
ˆu�U—\‘\Ñ"Ð$8Ð8Ñ	9ò]
ó _ó Fô]
r$   r×  )gH¯¼šò×z>)i,  r,   )Ar   rY  Údataclassesr   Útypingr   r   r   r   r    Útorch.utils.checkpointr   Úactivationsr
   Ú
file_utilsr   r   r   Úmodeling_outputsr   Úmodeling_utilsr   Úutilsr   r   Úutils.backbone_utilsr   Úconfiguration_zoedepthr   Ú
get_loggerr   Úloggerrë  r   ÚModuler'   r6   ru   r   ry   r£   r·   rÈ   rÊ   rà   rô   ÚjitÚscriptÚfloatro   r  r  r9  r=  rC  re  rz  r�  r•  r¾  rÅ  ÚZOEDEPTH_START_DOCSTRINGrê  r×  Ú__all__r#   r$   r%   ú<module>rÿ     sM  ðñ ã Ý !ß /Ó /ã Û Ý å !÷ñ õ
 5Ý -ß )Ý 1Ý 2ð 
ˆ×	Ñ	˜HÓ	%€ð #€ð ô? ;ó ?ó ð?ôBG˜bŸi™iô GôT˜bŸi™iô ô0# §¡ô #ô2;' "§)¡)ô ;'ô~" §¡ô "ôJ3$�2—9‘9ô 3$ôl))¨"¯)©)ô ))óXRô)0˜Ÿ™ô )0ôX@G¨B¯I©Iô @GôF5,˜rŸy™yô 5,ðp ‡�×Ññ-˜Uð -°ò -ó ð-ô&\,˜RŸY™Yô \,ô~Q, R§Y¡Yô Q,ôh˜Ÿ	™	ô ô6B §¡ô BôJ! b§i¡iô !ôH< b§i¡iô <ô~˜BŸI™Iô ô"J"°·±ô J"ôZ^¨¯	©	ô ^ôF*˜oô *ð0	Ð ðÐ ñ" ðð ó	ôy
Ð!8ó y
óðy
ðx (Ð)BÐ
C�r$   