Ë
    (täiœ=  ã                   ó†  — d dl Z d dlmZ d dlmc mZ ddlmZ ddlm	Z	 ddl
mZ  G d„ dej                  «      Z G d	„ d
ej                  «      Z G d„ dej                  «      Z G d„ dej                  «      Z G d„ dej                  «      Z	 	 	 dde j$                  de j$                  dz  dedede j$                  f
d„Zy)é    Né   )Ú	deprecateé   )ÚRMSNorm)Úupfirdn2d_nativec                   ó€   ‡ — e Zd ZdZ	 	 	 	 ddedededz  dedef
ˆ fd„Zd	ej                  d
ej                  fd„Z
ˆ xZS )ÚDownsample1Daÿ  A 1D downsampling layer with an optional convolution.

    Parameters:
        channels (`int`):
            number of channels in the inputs and outputs.
        use_conv (`bool`, default `False`):
            option to use a convolution.
        out_channels (`int`, optional):
            number of output channels. Defaults to `channels`.
        padding (`int`, default `1`):
            padding for the convolution.
        name (`str`, default `conv`):
            name of the downsampling 1D layer.
    NÚchannelsÚuse_convÚout_channelsÚpaddingÚnamec                 óN  •— t         ‰| �  «        || _        |xs || _        || _        || _        d}|| _        |r4t        j                  | j                  | j                  d||¬«      | _	        y | j                  | j                  k(  sJ ‚t        j                  ||¬«      | _	        y )Nr   é   ©Ústrider   ©Úkernel_sizer   )ÚsuperÚ__init__r
   r   r   r   r   ÚnnÚConv1dÚconvÚ	AvgPool1d)Úselfr
   r   r   r   r   r   Ú	__class__s          €úl/Volumes/fast/ai/experiments/MLX_z-image/.venv/lib/python3.12/site-packages/diffusers/models/downsampling.pyr   zDownsample1D.__init__(   sŽ   ø€ ô 	‰ÑÔØ ˆŒØ(Ò4¨HˆÔØ ˆŒØˆŒØˆØˆŒ	áÜŸ	™	 $§-¡-°×1BÑ1BÀAÈfÐ^eÔfˆD�Ià—=‘= D×$5Ñ$5Ò5Ð5Ð5ÜŸ™°ÀÔGˆD�Ió    ÚinputsÚreturnc                 ó`   — |j                   d   | j                  k(  sJ ‚| j                  |«      S )Nr   )Úshaper
   r   )r   r   s     r   ÚforwardzDownsample1D.forward>   s+   € Ø�|‰|˜A‰ $§-¡-Ò/Ð/Ð/Ø�y‰y˜Ó Ð r   )FNr   r   ©Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚintÚboolÚstrr   ÚtorchÚTensorr#   Ú__classcell__©r   s   @r   r	   r	      sp   ø„ ñð$ Ø#'ØØñHàðHð ðHð ˜D‘jð	Hð
 ðHð õHð,!˜eŸl™lð !¨u¯|©|÷ !r   r	   c                   óŠ   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 ddedededz  dedef
ˆ fd„Zd	ej                  d
ej                  fd„Z
ˆ xZS )ÚDownsample2Daÿ  A 2D downsampling layer with an optional convolution.

    Parameters:
        channels (`int`):
            number of channels in the inputs and outputs.
        use_conv (`bool`, default `False`):
            option to use a convolution.
        out_channels (`int`, optional):
            number of output channels. Defaults to `channels`.
        padding (`int`, default `1`):
            padding for the convolution.
        name (`str`, default `conv`):
            name of the downsampling 2D layer.
    Nr
   r   r   r   r   c                 ó0  •— t         ‰| �  «        || _        |xs || _        || _        || _        d}|| _        |dk(  rt        j                  |||	«      | _	        n0|dk(  rt        |||	«      | _	        n|€d | _	        nt        d|› �«      ‚|r0t        j                  | j                  | j                  ||||
¬«      }n2| j                  | j                  k(  sJ ‚t        j                  ||¬«      }|dk(  r|| _        || _        y |dk(  r|| _        y || _        y )	Nr   Úln_normÚrms_normzunknown norm_type: )r   r   r   Úbiasr   r   ÚConv2d_0)r   r   r
   r   r   r   r   r   Ú	LayerNormÚnormr   Ú
ValueErrorÚConv2dÚ	AvgPool2dr6   r   )r   r
   r   r   r   r   r   Ú	norm_typeÚepsÚelementwise_affiner5   r   r   r   s                €r   r   zDownsample2D.__init__S   s  ø€ ô 	‰ÑÔØ ˆŒØ(Ò4¨HˆÔØ ˆŒØˆŒØˆØˆŒ	à˜	Ò!ÜŸ™ X¨sÐ4FÓGˆD�IØ˜*Ò$Ü ¨#Ð/AÓBˆD�IØÐØˆD�IäÐ2°9°+Ð>Ó?Ð?áÜ—9‘9Ø—‘˜t×0Ñ0¸kÐRXÐbiÐptô‰Dð —=‘= D×$5Ñ$5Ò5Ð5Ð5Ü—<‘<¨F¸6ÔBˆDð �6Š>Ø ˆDŒMØˆD�IØ�ZÒØˆD�IàˆD�Ir   Úhidden_statesr    c                 óì  — t        |«      dkD  s|j                  dd «      �d}t        dd|«       |j                  d   | j                  k(  sJ ‚| j
                  �5| j                  |j                  dddd«      «      j                  dddd«      }| j                  r*| j                  dk(  rd}t        j                  ||d	d¬
«      }|j                  d   | j                  k(  sJ ‚| j                  |«      }|S )Nr   ÚscalezðThe `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`.z1.0.0r   r   r   ©r   r   r   r   Úconstant©ÚmodeÚvalue)ÚlenÚgetr   r"   r
   r8   Úpermuter   r   ÚFÚpadr   )r   r?   ÚargsÚkwargsÚdeprecation_messagerK   s         r   r#   zDownsample2D.forward‚   sê   € Üˆt‹9�qŠ=˜FŸJ™J w°Ó5ÐAð #UÐÜ�g˜wÐ(;Ô<Ø×"Ñ" 1Ñ%¨¯©Ò6Ð6Ð6à�9‰9Ð Ø ŸI™I m×&;Ñ&;¸A¸qÀ!ÀQÓ&GÓH×PÑPÐQRÐTUÐWXÐZ[Ó\ˆMà�=Š=˜TŸ\™\¨QÒ.ØˆCÜŸE™E -°¸:ÈQÔOˆMà×"Ñ" 1Ñ%¨¯©Ò6Ð6Ð6àŸ	™	 -Ó0ˆàÐr   )	FNr   r   r   NNNTr$   r/   s   @r   r1   r1   C   sz   ø„ ñð$ Ø#'ØØØØØØØñ-àð-ð ð-ð ˜D‘jð	-ð
 ð-ð õ-ð^ U§\¡\ð ÀuÇ|Á|÷ r   r1   c                   ó  ‡ — e Zd ZdZ	 	 	 	 ddedz  dedz  dedeeeeef   fˆ fd„Z	 	 	 	 ddej                  d	ej                  dz  d
ej                  dz  dede
dej                  fd„Zdej                  dej                  fd„Zˆ xZS )ÚFirDownsample2Da¼  A 2D FIR downsampling layer with an optional convolution.

    Parameters:
        channels (`int`):
            number of channels in the inputs and outputs.
        use_conv (`bool`, default `False`):
            option to use a convolution.
        out_channels (`int`, optional):
            number of output channels. Defaults to `channels`.
        fir_kernel (`tuple`, default `(1, 3, 3, 1)`):
            kernel for the FIR filter.
    Nr
   r   r   Ú
fir_kernelc                 óš   •— t         ‰| �  «        |r|n|}|rt        j                  ||ddd¬«      | _        || _        || _        || _        y )Nr   r   ©r   r   r   )r   r   r   r:   r6   rQ   r   r   )r   r
   r   r   rQ   r   s        €r   r   zFirDownsample2D.__init__¤   sL   ø€ ô 	‰ÑÔÙ'3‘|¸ˆÙÜŸI™I h°È!ÐTUÐ_`ÔaˆDŒMØ$ˆŒØ ˆŒØ(ˆÕr   r?   ÚweightÚkernelÚfactorÚgainr    c                 óÀ  — t        |t        «      r|dk\  sJ ‚|€dg|z  }t        j                  |t        j                  ¬«      }|j
                  dk(  rt        j                  ||«      }|t        j                  |«      z  }||z  }| j                  r€|j                  \  }}}}|j                  d   |z
  |dz
  z   }	||g}
t        |t        j                  ||j                  ¬«      |	dz   dz  |	dz  f¬«      }t        j                  |||
d¬«      }|S |j                  d   |z
  }	t        |t        j                  ||j                  ¬«      ||	dz   dz  |	dz  f¬«      }|S )	a"  Fused `Conv2d()` followed by `downsample_2d()`.
        Padding is performed only once at the beginning, not between the operations. The fused op is considerably more
        efficient than performing the same calculation using standard TensorFlow ops. It supports gradients of
        arbitrary order.

        Args:
            hidden_states (`torch.Tensor`):
                Input tensor of the shape `[N, C, H, W]` or `[N, H, W, C]`.
            weight (`torch.Tensor`, *optional*):
                Weight tensor of the shape `[filterH, filterW, inChannels, outChannels]`. Grouped convolution can be
                performed by `inChannels = x.shape[0] // numGroups`.
            kernel (`torch.Tensor`, *optional*):
                FIR filter of the shape `[firH, firW]` or `[firN]` (separable). The default is `[1] * factor`, which
                corresponds to average pooling.
            factor (`int`, *optional*, default to `2`):
                Integer downsampling factor.
            gain (`float`, *optional*, default to `1.0`):
                Scaling factor for signal magnitude.

        Returns:
            output (`torch.Tensor`):
                Tensor of the shape `[N, C, H // factor, W // factor]` or `[N, H // factor, W // factor, C]`, and same
                datatype as `x`.
        r   ©Údtyper   ©Údevicer   )rK   r   ©ÚdownrK   )Ú
isinstancer)   r,   ÚtensorÚfloat32ÚndimÚouterÚsumr   r"   r   r\   rJ   Úconv2d)r   r?   rT   rU   rV   rW   Ú_ÚconvHÚconvWÚ	pad_valueÚstride_valueÚupfirdn_inputÚoutputs                r   Ú_downsample_2dzFirDownsample2D._downsample_2d³   s_  € ôB ˜&¤#Ô&¨6°Qª;Ð6Ð6Øˆ>Ø�S˜6‘\ˆFô —‘˜f¬E¯M©MÔ:ˆØ�;‰;˜!ÒÜ—[‘[ ¨Ó0ˆFØ”%—)‘)˜FÓ#Ñ#ˆà˜$‘ˆà�=Š=Ø!'§¡ÑˆAˆq�%˜ØŸ™ a™¨6Ñ1°e¸a±iÑ@ˆIØ" FÐ+ˆLÜ,ØÜ—‘˜V¨M×,@Ñ,@ÔAØ !‘m¨Ñ)¨9¸©>Ð:ôˆMô
 —X‘X˜m¨V¸LÐRSÔTˆFð ˆð Ÿ™ Q™¨&Ñ0ˆIÜ%ØÜ—‘˜V¨M×,@Ñ,@ÔAØØ !‘m¨Ñ)¨9¸©>Ð:ô	ˆFð ˆr   c                 ó  — | j                   r_| j                  || j                  j                  | j                  ¬«      }|| j                  j
                  j                  dddd«      z   }|S | j                  || j                  d¬«      }|S )N)rT   rU   r   éÿÿÿÿr   )rU   rV   )r   rm   r6   rT   rQ   r5   Úreshape)r   r?   Údownsample_inputs      r   r#   zFirDownsample2D.forwardõ   s…   € Ø�=Š=Ø#×2Ñ2°=ÈÏÉ×I]ÑI]Ðfj×fuÑfuÐ2ÓvÐØ,¨t¯}©}×/AÑ/A×/IÑ/IÈ!ÈRÐQRÐTUÓ/VÑVˆMð Ðð !×/Ñ/°ÀdÇoÁoÐ^_Ð/Ó`ˆMàÐr   )NNF)r   r   r   r   )NNr   r   )r%   r&   r'   r(   r)   r*   Útupler   r,   r-   Úfloatrm   r#   r.   r/   s   @r   rP   rP   –   sá   ø„ ñð  $Ø#'ØØ0<ñ)à˜‘*ð)ð ˜D‘jð)ð ð	)ð
 ˜#˜s C¨Ð,Ñ-õ)ð$ '+Ø&*ØØñ@à—|‘|ð@ð —‘˜tÑ#ð@ð —‘˜tÑ#ð	@ð
 ð@ð ð@ð 
�‰ó@ðD U§\¡\ð °e·l±l÷ r   rP   c                   ób   ‡ — e Zd ZdZddefˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )ÚKDownsample2Dz‡A 2D K-downsampling layer.

    Parameters:
        pad_mode (`str`, *optional*, default to `"reflect"`): the padding mode to use.
    Úpad_modec                 óÖ   •— t         ‰| �  «        || _        t        j                  g d¢g«      }|j
                  d   dz  dz
  | _        | j                  d|j                  |z  d¬«       y )N)ç      À?ç      Ø?ry   rx   r   r   rU   F)Ú
persistent)	r   r   rv   r,   r`   r"   rK   Úregister_bufferÚT)r   rv   Ú	kernel_1dr   s      €r   r   zKDownsample2D.__init__  s_   ø€ Ü‰ÑÔØ ˆŒÜ—L‘LÒ">Ð!?Ó@ˆ	Ø—?‘? 1Ñ%¨Ñ*¨QÑ.ˆŒØ×Ñ˜X y§{¡{°YÑ'>È5ÐÕQr   r   r    c                 ó4  — t        j                  || j                  fdz  | j                  «      }|j                  |j                  d   |j                  d   | j
                  j                  d   | j
                  j                  d   g«      }t        j                  |j                  d   |j                  ¬«      }| j
                  j                  |«      d d d …f   j                  |j                  d   dd«      }||||f<   t        j                  ||d¬«      S )Né   r   r   r[   ro   r   )r   )rJ   rK   rv   Ú	new_zerosr"   rU   r,   Úaranger\   ÚtoÚexpandre   )r   r   rT   ÚindicesrU   s        r   r#   zKDownsample2D.forward  sã   € Ü—‘�v §¡˜{¨Q™°·±Ó>ˆØ×!Ñ!à—‘˜Q‘Ø—‘˜Q‘Ø—‘×!Ñ! !Ñ$Ø—‘×!Ñ! !Ñ$ð	ó
ˆô —,‘,˜vŸ|™|¨A™°v·}±}ÔEˆØ—‘—‘ Ó'¨ªa¨Ñ0×7Ñ7¸¿¹ÀQ¹ÈÈRÓPˆØ#)ˆˆw˜ÐÑ Ü�x‰x˜ ¨qÔ1Ð1r   )Úreflect)
r%   r&   r'   r(   r+   r   r,   r-   r#   r.   r/   s   @r   ru   ru      s1   ø„ ññR õ Rð2˜eŸl™lð 2¨u¯|©|÷ 2r   ru   c                   ó~   ‡ — e Zd ZdZ	 	 	 	 ddedededededefˆ fd„Zd	ej                  d
ej                  fd„Z	ˆ xZ
S )ÚCogVideoXDownsample3Da‹  
    A 3D Downsampling layer using in [CogVideoX]() by Tsinghua University & ZhipuAI

    Args:
        in_channels (`int`):
            Number of channels in the input image.
        out_channels (`int`):
            Number of channels produced by the convolution.
        kernel_size (`int`, defaults to `3`):
            Size of the convolving kernel.
        stride (`int`, defaults to `2`):
            Stride of the convolution.
        padding (`int`, defaults to `0`):
            Padding added to all four sides of the input.
        compress_time (`bool`, defaults to `False`):
            Whether or not to compress the time dimension.
    Úin_channelsr   r   r   r   Úcompress_timec                 ón   •— t         ‰| �  «        t        j                  |||||¬«      | _        || _        y )NrS   )r   r   r   r:   r   r‰   )r   rˆ   r   r   r   r   r‰   r   s          €r   r   zCogVideoXDownsample3D.__init__2  s2   ø€ ô 	‰ÑÔä—I‘I˜k¨<À[ÐY_ÐipÔqˆŒ	Ø*ˆÕr   Úxr    c                 óâ  — | j                   �r*|j                  \  }}}}}|j                  ddddd«      j                  ||z  |z  ||«      }|j                  d   dz  dk(  rŠ|d   |ddd …f   }}|j                  d   dkD  rt	        j
                  |dd¬	«      }t        j                  |d
   |gd¬«      }|j                  |||||j                  d   «      j                  ddddd«      }nMt	        j
                  |dd¬	«      }|j                  |||||j                  d   «      j                  ddddd«      }d}	t	        j                  ||	dd¬«      }|j                  \  }}}}}|j                  ddddd«      j                  ||z  |||«      }| j                  |«      }|j                  |||j                  d   |j                  d   |j                  d   «      j                  ddddd«      }|S )Nr   r   r   r   r   ro   ).r   .r   ).N)ÚdimrB   rC   rD   )
r‰   r"   rI   rp   rJ   Ú
avg_pool1dr,   ÚcatrK   r   )
r   r‹   Ú
batch_sizer
   ÚframesÚheightÚwidthÚx_firstÚx_restrK   s
             r   r#   zCogVideoXDownsample3D.forward@  sí  € Ø×ÓØ:;¿'¹'Ñ7ˆJ˜ &¨&°%ð —	‘	˜!˜Q  1 aÓ(×0Ñ0°¸fÑ1DÀuÑ1LÈhÐX^Ó_ˆAà�w‰w�r‰{˜Q‰ !Ò#Ø"# F¡)¨Q¨s°A±B¨w©Z˜�Ø—<‘< Ñ# aÒ'äŸ\™\¨&¸aÈÔJ�Fä—I‘I˜w yÑ1°6Ð:ÀÔC�à—I‘I˜j¨&°%¸À1Ç7Á7È2Á;ÓO×WÑWÐXYÐ[\Ð^_ÐabÐdeÓf‘ô —L‘L °¸!Ô<�à—I‘I˜j¨&°%¸À1Ç7Á7È2Á;ÓO×WÑWÐXYÐ[\Ð^_ÐabÐdeÓf�ð ˆÜ�E‰E�!�S˜z°Ô3ˆØ67·g±gÑ3ˆ
�H˜f f¨eà�I‰I�a˜˜A˜q !Ó$×,Ñ,¨Z¸&Ñ-@À(ÈFÐTYÓZˆØ�I‰I�a‹Lˆà�I‰I�j &¨!¯'©'°!©*°a·g±g¸a±jÀ!Ç'Á'È!Á*ÓM×UÑUÐVWÐYZÐ\]Ð_`ÐbcÓdˆØˆr   )r   r   r   F)r%   r&   r'   r(   r)   r*   r   r,   r-   r#   r.   r/   s   @r   r‡   r‡     sp   ø„ ñð, ØØØ#ñ+àð+ð ð+ð ð	+ð
 ð+ð ð+ð õ+ð˜Ÿ™ð ¨%¯,©,÷ r   r‡   r?   rU   rV   rW   r    c                 óž  — t        |t        «      r|dk\  sJ ‚|€dg|z  }t        j                  |t        j                  ¬«      }|j
                  dk(  rt        j                  ||«      }|t        j                  |«      z  }||z  }|j                  d   |z
  }t        | |j                  | j                  ¬«      ||dz   dz  |dz  f¬«      }|S )aE  Downsample2D a batch of 2D images with the given filter.
    Accepts a batch of 2D images of the shape `[N, C, H, W]` or `[N, H, W, C]` and downsamples each image with the
    given filter. The filter is normalized so that if the input pixels are constant, they will be scaled by the
    specified `gain`. Pixels outside the image are assumed to be zero, and the filter is padded with zeros so that its
    shape is a multiple of the downsampling factor.

    Args:
        hidden_states (`torch.Tensor`)
            Input tensor of the shape `[N, C, H, W]` or `[N, H, W, C]`.
        kernel (`torch.Tensor`, *optional*):
            FIR filter of the shape `[firH, firW]` or `[firN]` (separable). The default is `[1] * factor`, which
            corresponds to average pooling.
        factor (`int`, *optional*, default to `2`):
            Integer downsampling factor.
        gain (`float`, *optional*, default to `1.0`):
            Scaling factor for signal magnitude.

    Returns:
        output (`torch.Tensor`):
            Tensor of the shape `[N, C, H // factor, W // factor]`
    r   rY   r   r[   r   r]   )r_   r)   r,   r`   ra   rb   rc   rd   r"   r   r‚   r\   )r?   rU   rV   rW   ri   rl   s         r   Údownsample_2dr—   b  sÈ   € ô8 �fœcÔ" v°¢{Ð2Ð2Ø€~Ø��v‘ˆä�\‰\˜&¬¯©Ô6€FØ‡{�{�aÒÜ—‘˜V VÓ,ˆØ
Œe�i‰i˜ÓÑ€Fà�d‰]€FØ—‘˜Q‘ &Ñ(€IÜØØ�	‰	˜×-Ñ-ˆ	Ó.ØØ˜!‰m Ñ! 9°¡>Ð2ô	€Fð €Mr   )Nr   r   )r,   Útorch.nnr   Útorch.nn.functionalÚ
functionalrJ   Úutilsr   Únormalizationr   Ú
upsamplingr   ÚModuler	   r1   rP   ru   r‡   r-   r)   rs   r—   © r   r   Ú<module>r       sÅ   ðó Ý ß Ð å Ý "Ý (ô(!�2—9‘9ô (!ôVP�2—9‘9ô Pôff�b—i‘iô fôT2�B—I‘Iô 2ô<A˜BŸI™Iô AðL #'ØØñ	-Ø—<‘<ð-à�L‰L˜4Ñð-ð ð-ð ð	-ð
 ‡\�\ô-r   