Ë
    (täià ã                   óð  — d dl mZmZ d dlZd dlmZ d dlmc mZ ddl	m
Z
mZ ddlmZmZmZ ddlmZ ddlmZmZmZmZmZmZ dd	lmZmZmZ dd
lmZ ddlm Z m!Z!m"Z"m#Z#m$Z$  e«       rd dl%Z&ndZ& ejN                  e(«      Z) G d„ d«      Z* G d„ d«      Z+dejX                  dejZ                  de.de.fd„Z/e G d„ dejX                  «      «       Z0e G d„ dejX                  «      «       Z1e G d„ dejX                  «      «       Z2 G d„ dejX                  «      Z3e G d„ dejX                  «      «       Z4 G d„ d ejX                  «      Z5e G d!„ d"ejX                  «      «       Z6 G d#„ d$ejX                  «      Z7y)%é    )ÚAnyÚCallableNé   )Ú	deprecateÚlogging)Úis_torch_npu_availableÚis_torch_xla_availableÚis_xformers_available)Úmaybe_allow_in_graphé   )ÚGEGLUÚGELUÚApproximateGELUÚFP32SiLUÚLinearActivationÚSwiGLU)Ú	AttentionÚAttentionProcessorÚJointAttnProcessor2_0)ÚSinusoidalPositionalEmbedding)ÚAdaLayerNormÚAdaLayerNormContinuousÚAdaLayerNormZeroÚRMSNormÚSD35AdaLayerNormZeroXc                   óT   — e Zd Zedeeef   fd„«       Zdeeeef   z  fd„Zd„ Z	d„ Z
y)ÚAttentionMixinÚreturnc                 óÂ   ‡— i }dt         dt        j                  j                  dt        t         t
        f   fˆfd„Š| j                  «       D ]  \  }} ‰|||«       Œ |S )z¶
        Returns:
            `dict` of attention processors: A dictionary containing all attention processors used in the model with
            indexed by its weight name.
        ÚnameÚmoduleÚ
processorsc                 óš   •— t        |d«      r|j                  «       || › d�<   |j                  «       D ]  \  }} ‰| › d|› �||«       Œ |S )NÚget_processorú
.processorÚ.)Úhasattrr$   Únamed_children)r    r!   r"   Úsub_nameÚchildÚfn_recursive_add_processorss        €úi/Volumes/fast/ai/experiments/MLX_z-image/.venv/lib/python3.12/site-packages/diffusers/models/attention.pyr+   zCAttentionMixin.attn_processors.<locals>.fn_recursive_add_processors2   s]   ø€ Ü�v˜Ô/Ø28×2FÑ2FÓ2H�
˜d˜V :Ð.Ñ/à#)×#8Ñ#8Ö#:‘�˜%Ù+¨t¨f°A°h°ZÐ,@À%ÈÕTð $;ð Ðó    )ÚstrÚtorchÚnnÚModuleÚdictr   r(   )Úselfr"   r    r!   r+   s       @r,   Úattn_processorszAttentionMixin.attn_processors(   sf   ø€ ð ˆ
ð	¬cð 	¼5¿8¹8¿?¹?ð 	ÔX\Ô]`ÔbtÐ]tÑXuõ 	ð !×/Ñ/Ö1‰LˆD�&Ù'¨¨f°jÕAð 2ð Ðr-   Ú	processorc           	      óT  ‡— t        | j                  j                  «       «      }t        |t        «      r,t        |«      |k7  rt        dt        |«      › d|› d|› d�«      ‚dt        dt        j                  j                  fˆfd„Š| j                  «       D ]  \  }} ‰|||«       Œ y)	a4  
        Sets the attention processor to use to compute attention.

        Parameters:
            processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
                The instantiated processor class or a dictionary of processor classes that will be set as the processor
                for **all** `Attention` layers.

                If `processor` is a dict, the key needs to define the path to the corresponding cross attention
                processor. This is strongly recommended when setting trainable attention processors.

        z>A dict of processors was passed, but the number of processors z0 does not match the number of attention layers: z. Please make sure to pass z processor classes.r    r!   c                 óö   •— t        |d«      rEt        |t        «      s|j                  |«       n#|j                  |j	                  | › d�«      «       |j                  «       D ]  \  }} ‰| › d|› �||«       Œ y )NÚset_processorr%   r&   )r'   Ú
isinstancer2   r8   Úpopr(   )r    r!   r5   r)   r*   Úfn_recursive_attn_processors        €r,   r;   zFAttentionMixin.set_attn_processor.<locals>.fn_recursive_attn_processorU   sq   ø€ Ü�v˜Ô/Ü! )¬TÔ2Ø×(Ñ(¨Õ3à×(Ñ(¨¯©¸$¸¸zÐ7JÓ)KÔLà#)×#8Ñ#8Ö#:‘�˜%Ù+¨t¨f°A°h°ZÐ,@À%ÈÕSñ $;r-   N)Úlenr4   Úkeysr9   r2   Ú
ValueErrorr.   r/   r0   r1   r(   )r3   r5   Úcountr    r!   r;   s        @r,   Úset_attn_processorz!AttentionMixin.set_attn_processor@   s¯   ø€ ô �D×(Ñ(×-Ñ-Ó/Ó0ˆä�i¤Ô&¬3¨y«>¸UÒ+BÜØPÔQTÐU^ÓQ_ÐP`ð a0Ø05¨wÐ6QÐRWÐQXÐXkðmóð ð
	T¬cð 	T¼5¿8¹8¿?¹?õ 	Tð !×/Ñ/Ö1‰LˆD�&Ù'¨¨f°iÕ@ñ 2r-   c                 ó&  — | j                   j                  «       D ]1  \  }}dt        |j                  j                  «      v sŒ(t        d«      ‚ | j                  «       D ]0  }t        |t        «      sŒ|j                  sŒ!|j                  «        Œ2 y)zÛ
        Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value)
        are fused. For cross-attention modules, key and value projection matrices are fused.
        ÚAddedzQ`fuse_qkv_projections()` is not supported for models having added KV projections.N)r4   Úitemsr.   Ú	__class__Ú__name__r>   Úmodulesr9   ÚAttentionModuleMixinÚ_supports_qkv_fusionÚfuse_projections)r3   Ú_Úattn_processorr!   s       r,   Úfuse_qkv_projectionsz#AttentionMixin.fuse_qkv_projectionsb   sx   € ð
 "&×!5Ñ!5×!;Ñ!;Ö!=ÑˆAˆ~Øœ#˜n×6Ñ6×?Ñ?Ó@Ò@Ü Ð!tÓuÐuð ">ð —l‘l–nˆFÜ˜&Ô"6Õ7¸F×<WÓ<WØ×'Ñ'Õ)ñ %r-   c                 óŠ   — | j                  «       D ]0  }t        |t        «      sŒ|j                  sŒ!|j	                  «        Œ2 y)um   Disables the fused QKV projection if enabled.

        > [!WARNING] > This API is ðŸ§ª experimental.
        N)rF   r9   rG   rH   Úunfuse_projections)r3   r!   s     r,   Úunfuse_qkv_projectionsz%AttentionMixin.unfuse_qkv_projectionso   s3   € ð
 —l‘l–nˆFÜ˜&Ô"6Õ7¸F×<WÓ<WØ×)Ñ)Õ+ñ %r-   N)rE   Ú
__module__Ú__qualname__Úpropertyr2   r.   r   r4   r@   rL   rO   © r-   r,   r   r   '   sP   „ Øð  cÐ+=Ð&=Ñ!>ò ó ðð. AÐ,>ÀÀcÐK]ÐF]ÑA^Ñ,^ó  AòD*ó,r-   r   c                   ó|  — e Zd ZdZg ZdZdZdeddfd„Zd&de	ddfd	„Z
d
efd„Zde	ddfd„Z	 	 d'de	deedz  df   dz  ddfd„Z	 d(de	dedz  ddfd„Z ej&                  «       d„ «       Z ej&                  «       d„ «       Zdeddfd„Zdej0                  dej0                  fd„Zd)dej0                  dedej0                  fd„Z	 d(dej0                  dej0                  dej0                  dz  dej0                  fd „Z	 d)dej0                  d!ed"ededej0                  f
d#„Zd$ej0                  dej0                  fd%„Zy)*rG   NTFr5   r   c                 óN  — t        | d«      r’t        | j                  t        j                  j
                  «      rdt        |t        j                  j
                  «      s@t        j                  d| j                  › d|› �«       | j                  j                  d«       || _        y)z�
        Set the attention processor to use.

        Args:
            processor (`AttnProcessor`):
                The attention processor to use.
        r5   z-You are removing possibly trained weights of z with N)
r'   r9   r5   r/   r0   r1   ÚloggerÚinfoÚ_modulesr:   )r3   r5   s     r,   r8   z"AttentionModuleMixin.set_processor   sq   € ô �D˜+Ô&Ü˜4Ÿ>™>¬5¯8©8¯?©?Ô;Ü˜y¬%¯(©(¯/©/Ô:ä�K‰KÐGÈÏÉÐGWÐW]Ð^gÐ]hÐiÔjØ�M‰M×Ñ˜kÔ*à"ˆ�r-   Úreturn_deprecated_lorar   c                 ó    — |s| j                   S y)a7  
        Get the attention processor in use.

        Args:
            return_deprecated_lora (`bool`, *optional*, defaults to `False`):
                Set to `True` to return the deprecated LoRA attention processor.

        Returns:
            "AttentionProcessor": The attention processor in use.
        N)r5   )r3   rY   s     r,   r$   z"AttentionModuleMixin.get_processor“   s   € ñ &Ø—>‘>Ð!ð &r-   Úbackendc                 ó  — ddl m} |j                  j                  «       D �ch c]  }|j                  ’Œ }}||vr!t        d|›d�dj                  |«      z   «      ‚ ||j                  «       «      }|| j                  _	        y c c}w )Nr   )ÚAttentionBackendNamez	`backend=z ` must be one of the following: z, )
Úattention_dispatchr]   Ú__members__ÚvaluesÚvaluer>   ÚjoinÚlowerr5   Ú_attention_backend)r3   r[   r]   ÚxÚavailable_backendss        r,   Úset_attention_backendz*AttentionModuleMixin.set_attention_backend¡   s~   € Ý<à/C×/OÑ/O×/VÑ/VÔ/XÓYÑ/X¨!˜aŸg›gÐ/XÐÐYØÐ,Ñ,Ü˜z  
Ð*JÐKÈdÏiÉiÐXjÓNkÑkÓlÐlá& w§}¡}£Ó7ˆØ,3ˆ�‰Õ)ùò Zs   £BÚuse_npu_flash_attentionc                 óT   — |rt        «       st        d«      ‚| j                  d«       y)z¹
        Set whether to use NPU flash attention from `torch_npu` or not.

        Args:
            use_npu_flash_attention (`bool`): Whether to use NPU flash attention or not.
        ztorch_npu is not availableÚ_native_npuN)r   ÚImportErrorrg   )r3   rh   s     r,   Úset_use_npu_flash_attentionz0AttentionModuleMixin.set_use_npu_flash_attention«   s'   € ñ #Ü)Ô+Ü!Ð">Ó?Ð?à×"Ñ" =Õ1r-   Úuse_xla_flash_attentionÚpartition_spec.c                 óT   — |rt        «       st        d«      ‚| j                  d«       y)aÝ  
        Set whether to use XLA flash attention from `torch_xla` or not.

        Args:
            use_xla_flash_attention (`bool`):
                Whether to use pallas flash attention kernel from `torch_xla` or not.
            partition_spec (`tuple[]`, *optional*):
                Specify the partition specification if using SPMD. Otherwise None.
            is_flux (`bool`, *optional*, defaults to `False`):
                Whether the model is a Flux model.
        ztorch_xla is not availableÚ_native_xlaN)r	   rk   rg   )r3   rm   rn   Úis_fluxs       r,   Úset_use_xla_flash_attentionz0AttentionModuleMixin.set_use_xla_flash_attention¹   s'   € ñ" #Ü)Ô+Ü!Ð">Ó?Ð?à×"Ñ" =Õ1r-   Ú'use_memory_efficient_attention_xformersÚattention_opc                 óˆ  — |r­t        «       st        dd¬«      ‚t        j                  j	                  «       st        d«      ‚	 t        «       rPd}|�|\  }}|j                  ^}}t        j                  dd|¬«      }t        j                  j                  |||«      }| j                  d«       yy# t        $ r}|‚d}~ww xY w)	a¸  
        Set whether to use memory efficient attention from `xformers` or not.

        Args:
            use_memory_efficient_attention_xformers (`bool`):
                Whether to use memory efficient attention from `xformers` or not.
            attention_op (`Callable`, *optional*):
                The attention operation to use. Defaults to `None` which uses the default attention operation from
                `xformers`.
        zeRefer to https://github.com/facebookresearch/xformers for more information on how to install xformersÚxformers)r    zvtorch.cuda.is_available() should be True but is False. xformers' memory efficient attention is only available for GPU N)r   r   é(   Úcuda©ÚdeviceÚdtype)r
   ÚModuleNotFoundErrorr/   rx   Úis_availabler>   ÚSUPPORTED_DTYPESÚrandnÚxopsÚopsÚmemory_efficient_attentionÚ	Exceptionrg   )	r3   rs   rt   r{   Úop_fwÚop_bwrJ   ÚqÚes	            r,   Ú+set_use_memory_efficient_attention_xformersz@AttentionModuleMixin.set_use_memory_efficient_attention_xformersÐ   sÈ   € ñ 3Ü(Ô*Ü)Ø{Ø#ôð ô —Z‘Z×,Ñ,Ô.Ü ð/óð ð

ä,Ô.Ø $˜Ø'Ð3Ø+7™L˜E 5Ø(-×(>Ñ(>˜I˜E AÜ!ŸK™K¨
¸6ÈÔO˜Ü ŸH™H×?Ñ?ÀÀ1ÀaÓH˜ð ×*Ñ*¨:Õ6ð1 3øô* !ò Ø�Gûðús   ÁAB1 Â1	CÂ:B<Â<Cc                 ó’
  — | j                   s-t        j                  | j                  j                  › d�«       yt        | dd«      ry| j                  j                  j                  j                  }| j                  j                  j                  j                  }t        | d«      �r`| j                  �rSt        j                  | j                  j                  j                  | j                   j                  j                  g«      }|j"                  d   }|j"                  d   }t%        j&                  ||| j(                  ||¬«      | _        | j*                  j                  j-                  |«       t        | d	«      �r| j(                  �rt        j                  | j                  j.                  j                  | j                   j.                  j                  g«      }| j*                  j.                  j-                  |«       �n�t        j                  | j                  j                  j                  | j                  j                  j                  | j                   j                  j                  g«      }|j"                  d   }|j"                  d   }t%        j&                  ||| j(                  ||¬«      | _        | j0                  j                  j-                  |«       t        | d	«      r£| j(                  r—t        j                  | j                  j.                  j                  | j                  j.                  j                  | j                   j.                  j                  g«      }| j0                  j.                  j-                  |«       t        | d
d«      ���t        | dd«      ���t        | dd«      ���t        j                  | j2                  j                  j                  | j4                  j                  j                  | j6                  j                  j                  g«      }|j"                  d   }|j"                  d   }t%        j&                  ||| j8                  ||¬«      | _        | j:                  j                  j-                  |«       | j8                  r—t        j                  | j2                  j.                  j                  | j4                  j.                  j                  | j6                  j.                  j                  g«      }| j:                  j.                  j-                  |«       d| _        y)ze
        Fuse the query, key, and value projections into a single projection for efficiency.
        zK does not support fusing QKV projections, so `fuse_projections` will no-op.NÚfused_projectionsFÚis_cross_attentionr   r   )Úbiasrz   r{   Úuse_biasÚ
add_q_projÚ
add_k_projÚ
add_v_projT)rH   rV   ÚdebugrD   rE   ÚgetattrÚto_qÚweightÚdatarz   r{   r'   r‹   r/   ÚcatÚto_kÚto_vÚshaper0   ÚLinearr�   Úto_kvÚcopy_rŒ   Úto_qkvrŽ   r�   r�   Úadded_proj_biasÚto_added_qkvrŠ   )r3   rz   r{   Úconcatenated_weightsÚin_featuresÚout_featuresÚconcatenated_biass          r,   rI   z%AttentionModuleMixin.fuse_projections÷   sÒ  € ð ×(Ò(Ü�L‰LØ—>‘>×*Ñ*Ð+Ð+vÐwôð ô �4Ð,¨eÔ4Øà—‘×!Ñ!×&Ñ&×-Ñ-ˆØ—	‘	× Ñ ×%Ñ%×+Ñ+ˆä�4Ð-Õ.°4×3JÓ3Jä#(§9¡9¨d¯i©i×.>Ñ.>×.CÑ.CÀTÇYÁY×EUÑEU×EZÑEZÐ-[Ó#\Ð Ø.×4Ñ4°QÑ7ˆKØ/×5Ñ5°aÑ8ˆLäŸ™ ;°À4Ç=Á=ÐY_ÐglÔmˆDŒJØ�J‰J×Ñ×#Ñ#Ð$8Ô9Ü�t˜ZÕ(¨T¯]«]Ü$)§I¡I¨t¯y©y¯~©~×/BÑ/BÀDÇIÁIÇNÁN×DWÑDWÐ.XÓ$YÐ!Ø—
‘
—‘×%Ñ%Ð&7Ö8ô $)§9¡9¨d¯i©i×.>Ñ.>×.CÑ.CÀTÇYÁY×EUÑEU×EZÑEZÐ\`×\eÑ\e×\lÑ\l×\qÑ\qÐ-rÓ#sÐ Ø.×4Ñ4°QÑ7ˆKØ/×5Ñ5°aÑ8ˆLäŸ)™) K°ÀDÇMÁMÐZ`ÐhmÔnˆDŒKØ�K‰K×Ñ×$Ñ$Ð%9Ô:Ü�t˜ZÔ(¨T¯]ª]Ü$)§I¡I¨t¯y©y¯~©~×/BÑ/BÀDÇIÁIÇNÁN×DWÑDWÐY]×YbÑYb×YgÑYg×YlÑYlÐ.mÓ$nÐ!Ø—‘× Ñ ×&Ñ&Ð'8Ô9ô �D˜,¨Ó-Ñ9Ü˜˜l¨DÓ1Ñ=Ü˜˜l¨DÓ1Ñ=ä#(§9¡9Ø—‘×'Ñ'×,Ñ,¨d¯o©o×.DÑ.D×.IÑ.IÈ4Ï?É?×KaÑKa×KfÑKfÐgó$Ð ð /×4Ñ4°QÑ7ˆKØ/×5Ñ5°aÑ8ˆLä "§	¡	Ø˜\°×0DÑ0DÈVÐ[`ô!ˆDÔð ×Ñ×$Ñ$×*Ñ*Ð+?Ô@Ø×#Ò#Ü$)§I¡IØ—_‘_×)Ñ)×.Ñ.°·±×0DÑ0D×0IÑ0IÈ4Ï?É?×K_ÑK_×KdÑKdÐeó%Ð!ð ×!Ñ!×&Ñ&×,Ñ,Ð->Ô?à!%ˆÕr-   c                 óØ   — | j                   syt        | dd«      syt        | d«      rt        | d«       t        | d«      rt        | d«       t        | d«      rt        | d«       d| _        y)z\
        Unfuse the query, key, and value projections back to separate projections.
        NrŠ   Fr�   r›   rŸ   )rH   r’   r'   ÚdelattrrŠ   )r3   s    r,   rN   z'AttentionModuleMixin.unfuse_projections:  sh   € ð ×(Ò(Øô �tÐ0°%Ô8Øô �4˜Ô"Ü�D˜(Ô#ä�4˜Ô!Ü�D˜'Ô"ä�4˜Ô(Ü�D˜.Ô)à!&ˆÕr-   Ú
slice_sizec                 óæ   — t        | d«      r-|�+|| j                  kD  rt        d|› d| j                  › d�«      ‚d}|�| j                  d«      }|€| j	                  «       }| j                  |«       y)z¨
        Set the slice size for attention computation.

        Args:
            slice_size (`int`):
                The slice size for attention computation.
        Úsliceable_head_dimNzslice_size z has to be smaller or equal to r&   Úsliced)r'   r¨   r>   Ú_get_compatible_processorÚdefault_processor_clsr8   )r3   r¦   r5   s      r,   Úset_attention_slicez(AttentionModuleMixin.set_attention_sliceT  s‡   € ô �4Ð-Ô.°:Ð3IÈjÐ[_×[rÑ[rÒNrÜ˜{¨:¨,Ð6UÐVZ×VmÑVmÐUnÐnoÐpÓqÐqàˆ	ð Ð!Ø×6Ñ6°xÓ@ˆIð ÐØ×2Ñ2Ó4ˆIà×Ñ˜9Õ%r-   Útensorc                 óÂ   — | j                   }|j                  \  }}}|j                  ||z  |||«      }|j                  dddd«      j                  ||z  |||z  «      }|S )a  
        Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size // heads, seq_len, dim * heads]`.

        Args:
            tensor (`torch.Tensor`): The tensor to reshape.

        Returns:
            `torch.Tensor`: The reshaped tensor.
        r   r   r   é   )Úheadsr™   ÚreshapeÚpermute)r3   r­   Ú	head_sizeÚ
batch_sizeÚseq_lenÚdims         r,   Úbatch_to_head_dimz&AttentionModuleMixin.batch_to_head_dimk  sj   € ð —J‘Jˆ	Ø#)§<¡<Ñ ˆ
�G˜SØ—‘ 
¨iÑ 7¸ÀGÈSÓQˆØ—‘  1 a¨Ó+×3Ñ3°JÀ)Ñ4KÈWÐVYÐ\eÑVeÓfˆØˆr-   Úout_dimc                 ó"  — | j                   }|j                  dk(  r|j                  \  }}}d}n|j                  \  }}}}|j                  |||z  |||z  «      }|j	                  dddd«      }|dk(  r|j                  ||z  ||z  ||z  «      }|S )a5  
        Reshape the tensor for multi-head attention processing.

        Args:
            tensor (`torch.Tensor`): The tensor to reshape.
            out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor.

        Returns:
            `torch.Tensor`: The reshaped tensor.
        r¯   r   r   r   )r°   Úndimr™   r±   r²   )r3   r­   r¸   r³   r´   rµ   r¶   Ú	extra_dims           r,   Úhead_to_batch_dimz&AttentionModuleMixin.head_to_batch_dim{  s¢   € ð —J‘Jˆ	Ø�;‰;˜!ÒØ'-§|¡|Ñ$ˆJ˜ Ø‰Ià28·,±,Ñ/ˆJ˜	 7¨CØ—‘ 
¨G°iÑ,?ÀÈCÐS\ÑL\Ó]ˆØ—‘  1 a¨Ó+ˆà�aŠ<Ø—^‘^ J°Ñ$:¸GÀiÑ<OÐQTÐXaÑQaÓbˆFàˆr-   ÚqueryÚkeyÚattention_maskc                 ó  — |j                   }| j                  r |j                  «       }|j                  «       }|€Xt        j                  |j
                  d   |j
                  d   |j
                  d   |j                   |j                  ¬«      }d}n|}d}t        j                  |||j                  dd«      || j                  ¬«      }~| j                  r|j                  «       }|j                  d¬«      }~|j                  |«      }|S )aL  
        Compute the attention scores.

        Args:
            query (`torch.Tensor`): The query tensor.
            key (`torch.Tensor`): The key tensor.
            attention_mask (`torch.Tensor`, *optional*): The attention mask to use.

        Returns:
            `torch.Tensor`: The attention probabilities/scores.
        r   r   ©r{   rz   éÿÿÿÿéþÿÿÿ)ÚbetaÚalpha©r¶   )r{   Úupcast_attentionÚfloatr/   Úemptyr™   rz   ÚbaddbmmÚ	transposeÚscaleÚupcast_softmaxÚsoftmaxÚto)	r3   r½   r¾   r¿   r{   Úbaddbmm_inputrÄ   Úattention_scoresÚattention_probss	            r,   Úget_attention_scoresz)AttentionModuleMixin.get_attention_scores”  sõ   € ð —‘ˆØ× Ò Ø—K‘K“MˆEØ—)‘)“+ˆCàÐ!Ü!ŸK™KØ—‘˜A‘ §¡¨A¡°·	±	¸!±ÀEÇKÁKÐX]×XdÑXdôˆMð ‰Dà*ˆMØˆDä Ÿ=™=ØØØ�M‰M˜"˜bÓ!ØØ—*‘*ô
Ðð à×ÒØ/×5Ñ5Ó7Ðà*×2Ñ2°rÐ2Ó:ˆØà)×,Ñ,¨UÓ3ˆàÐr-   Útarget_lengthr´   c                 ó.  — | j                   }|€|S |j                  d   }||k7  r˜|j                  j                  dk(  re|j                  d   |j                  d   |f}t	        j
                  ||j                  |j                  ¬«      }t	        j                  ||gd¬«      }nt        j                  |d|fd¬	«      }|d
k(  r*|j                  d   ||z  k  r|j                  |d¬«      }|S |dk(  r$|j                  d«      }|j                  |d¬«      }|S )aÚ  
        Prepare the attention mask for the attention computation.

        Args:
            attention_mask (`torch.Tensor`): The attention mask to prepare.
            target_length (`int`): The target length of the attention mask.
            batch_size (`int`): The batch size for repeating the attention mask.
            out_dim (`int`, *optional*, defaults to `3`): Output dimension.

        Returns:
            `torch.Tensor`: The prepared attention mask.
        rÂ   Úmpsr   r   rÁ   r   rÆ   ç        )ra   r¯   é   )r°   r™   rz   Útyper/   Úzerosr{   r–   ÚFÚpadÚrepeat_interleaveÚ	unsqueeze)	r3   r¿   rÔ   r´   r¸   r³   Úcurrent_lengthÚpadding_shapeÚpaddings	            r,   Úprepare_attention_maskz+AttentionModuleMixin.prepare_attention_maskÃ  s)  € ð —J‘Jˆ	ØÐ!Ø!Ð!à,×2Ñ2°2Ñ6ˆØ˜]Ò*Ø×$Ñ$×)Ñ)¨UÒ2ð "0×!5Ñ!5°aÑ!8¸.×:NÑ:NÈqÑ:QÐS`Ð a�ÜŸ+™+ m¸>×;OÑ;OÐXf×XmÑXmÔn�Ü!&§¡¨N¸GÐ+DÈ!Ô!L‘ô "#§¡ ~¸¸=Ð7IÐQTÔ!U�à�aŠ<Ø×#Ñ# AÑ&¨°iÑ)?Ò?Ø!/×!AÑ!AÀ)ÐQRÐ!AÓ!S�ð
 Ðð	 ˜Š\Ø+×5Ñ5°aÓ8ˆNØ+×=Ñ=¸iÈQÐ=ÓOˆNàÐr-   Úencoder_hidden_statesc                 óP  — | j                   €J d«       ‚t        | j                   t        j                  «      r| j                  |«      }|S t        | j                   t        j                  «      r7|j                  dd«      }| j                  |«      }|j                  dd«      }|S J ‚)zë
        Normalize the encoder hidden states.

        Args:
            encoder_hidden_states (`torch.Tensor`): Hidden states of the encoder.

        Returns:
            `torch.Tensor`: The normalized encoder hidden states.
        zGself.norm_cross must be defined to call self.norm_encoder_hidden_statesr   r   )Ú
norm_crossr9   r0   Ú	LayerNormÚ	GroupNormrË   )r3   rã   s     r,   Únorm_encoder_hidden_statesz/AttentionModuleMixin.norm_encoder_hidden_statesî  sŸ   € ð �‰Ð*ÐuÐ,uÓuÐ*Ü�d—o‘o¤r§|¡|Ô4Ø$(§O¡OÐ4IÓ$JÐ!ð %Ð$ô ˜Ÿ™¬¯©Ô6ð %:×$CÑ$CÀAÀqÓ$IÐ!Ø$(§O¡OÐ4IÓ$JÐ!Ø$9×$CÑ$CÀAÀqÓ$IÐ!ð %Ð$ð �5r-   )F)NF©N)r¯   )rE   rP   rQ   Ú_default_processor_clsÚ_available_processorsrH   rŠ   r   r8   Úboolr$   r.   rg   rl   Útuplerr   r   rˆ   r/   Úno_gradrI   rN   Úintr¬   ÚTensorr·   r¼   rÓ   râ   rè   rS   r-   r,   rG   rG   y   sï  „ Ø!ÐØÐØÐØÐð#Ð'9ð #¸dó #ñ("°Dð "ÐEYó "ð4¨Só 4ð2À4ð 2ÈDó 2ð" 9=Øñ	2à!%ð2ð ˜c D™j¨#˜oÑ.°Ñ5ð2ð
 
ó2ð0 ^bñ%7Ø7;ð%7ØKSÐVZÉ?ð%7à	ó%7ðN €U‡]�]ƒ_ñ@&ó ð@&ðD €U‡]�]ƒ_ñ'ó ð'ð2&¨cð &°dó &ð.¨¯©ð ¸¿¹ó ñ ¨¯©ð ¸sð È5Ï<É<ó ð4 ]añ-Ø—\‘\ð-Ø(-¯©ð-ØFKÇlÁlÐUYÑFYð-à	�‰ó-ð` abñ)Ø#Ÿl™lð)Ø;>ð)ØLOð)ØZ]ð)à	�‰ó)ðV%ÀÇÁð %ÐQV×Q]ÑQ]ô %r-   rG   ÚffÚhidden_statesÚ	chunk_dimÚ
chunk_sizec                 ó  — |j                   |   |z  dk7  rt        d|j                   |   › d|› d�«      ‚|j                   |   |z  }t        j                  |j	                  ||¬«      D �cg c]
  } | |«      ‘Œ c}|¬«      }|S c c}w )Nr   z)`hidden_states` dimension to be chunked: z$ has to be divisible by chunk size: z[. Make sure to set an appropriate `chunk_size` when calling `unet.enable_forward_chunking`.rÆ   )r™   r>   r/   r–   Úchunk)rñ   rò   ró   rô   Ú
num_chunksÚ	hid_sliceÚ	ff_outputs          r,   Ú_chunked_feed_forwardrú   
  s¶   € à×Ñ˜9Ñ%¨
Ñ2°aÒ7ÜØ7¸×8KÑ8KÈIÑ8VÐ7WÐW{ð  }Gð  |Hð  Hcð  dó
ð 	
ð ×$Ñ$ YÑ/°:Ñ=€JÜ—	‘	Ø(5×(;Ñ(;¸JÈIÐ(;Ô(VÓWÑ(V˜9‰ˆI�Ð(VÑWØô€Ið Ðùò 	Xs   Á+Bc                   ó„   ‡ — e Zd ZdZdedededefˆ fd„Zdej                  dej                  d	ej                  fd
„Zˆ xZ	S )ÚGatedSelfAttentionDenseat  
    A gated self-attention dense layer that combines visual features and object features.

    Parameters:
        query_dim (`int`): The number of channels in the query.
        context_dim (`int`): The number of channels in the context.
        n_heads (`int`): The number of heads to use for attention.
        d_head (`int`): The number of channels in each head.
    Ú	query_dimÚcontext_dimÚn_headsÚd_headc                 óø  •— t         ‰| �  «        t        j                  ||«      | _        t        |||¬«      | _        t        |d¬«      | _        t        j                  |«      | _
        t        j                  |«      | _        | j                  dt        j                  t        j                  d«      «      «       | j                  dt        j                  t        j                  d«      «      «       d| _        y )N)rý   r°   Údim_headÚgeglu©Úactivation_fnÚ
alpha_attnr×   Úalpha_denseT)ÚsuperÚ__init__r0   rš   Úlinearr   ÚattnÚFeedForwardrñ   ræ   Únorm1Únorm2Úregister_parameterÚ	Parameterr/   r­   Úenabled)r3   rý   rþ   rÿ   r   rD   s        €r,   r	  z GatedSelfAttentionDense.__init__%  s«   ø€ Ü‰ÑÔô —i‘i ¨YÓ7ˆŒä¨	¸È6ÔRˆŒ	Ü˜i°wÔ?ˆŒä—\‘\ )Ó,ˆŒ
Ü—\‘\ )Ó,ˆŒ
à×Ñ ¬b¯l©l¼5¿<¹<ÈÓ;LÓ.MÔNØ×Ñ ¬r¯|©|¼E¿L¹LÈÓ<MÓ/NÔOàˆ�r-   re   Úobjsr   c           
      ó   — | j                   s|S |j                  d   }| j                  |«      }|| j                  j	                  «       | j                  | j                  t        j                  ||gd¬«      «      «      d d …d |…d d …f   z  z   }|| j                  j	                  «       | j                  | j                  |«      «      z  z   }|S )Nr   rÆ   )r  r™   r
  r  Útanhr  r  r/   r–   r  rñ   r  )r3   re   r  Ún_visuals       r,   ÚforwardzGatedSelfAttentionDense.forward6  s°   € Ø�|Š|ØˆHà—7‘7˜1‘:ˆØ�{‰{˜4Ó ˆà�—‘×$Ñ$Ó&¨¯©°4·:±:¼e¿i¹iÈÈDÈ	ÐWXÔ>YÓ3ZÓ)[Ò\]Ð_hÐ`hÐ_hÒjkÐ\kÑ)lÑlÑlˆØ�× Ñ ×%Ñ%Ó'¨$¯'©'°$·*±*¸Q³-Ó*@Ñ@Ñ@ˆàˆr-   )
rE   rP   rQ   Ú__doc__rï   r	  r/   rð   r  Ú__classcell__©rD   s   @r,   rü   rü     sO   ø„ ñð #ð °Cð À#ð Èsõ ð"
˜Ÿ™ð 
¨U¯\©\ð 
¸e¿l¹l÷ 
r-   rü   c                   ó   ‡ — e Zd ZdZ	 	 	 ddedededededz  defˆ fd	„Zdd
edz  defd„Z	 dde	j                  de	j                  de	j                  deeef   dz  dee	j                  e	j                  f   f
d„Zˆ xZS )ÚJointTransformerBlocka,  
    A Transformer block following the MMDiT architecture, introduced in Stable Diffusion 3.

    Reference: https://huggingface.co/papers/2403.03206

    Parameters:
        dim (`int`): The number of channels in the input and output.
        num_attention_heads (`int`): The number of heads to use for multi-head attention.
        attention_head_dim (`int`): The number of channels in each head.
        context_pre_only (`bool`): Boolean to determine if we should add some blocks associated with the
            processing of `context` conditions.
    Nr¶   Únum_attention_headsÚattention_head_dimÚcontext_pre_onlyÚqk_normÚuse_dual_attentionc                 óØ  •— t         ‰	| �  «        || _        || _        |rdnd}|rt	        |«      | _        nt        |«      | _        |dk(  rt        ||dddd¬«      | _        n%|dk(  rt        |«      | _        nt        d|› d	�«      ‚t        t        d
«      rt        «       }nt        d«      ‚t        |d |||||d||d¬«      | _        |rt        |d |||d||d¬«	      | _        nd | _        t!        j"                  |dd¬«      | _        t'        ||d¬«      | _        |s1t!        j"                  |dd¬«      | _        t'        ||d¬«      | _        nd | _        d | _        d | _        d| _        y )NÚada_norm_continousÚada_norm_zeroFç�íµ ÷Æ°>TÚ
layer_norm)Úelementwise_affineÚepsrŒ   Ú	norm_typezUnknown context_norm_type: z>, currently only support `ada_norm_continous`, `ada_norm_zero`Úscaled_dot_product_attentionzYThe current PyTorch version does not support the `scaled_dot_product_attention` function.)rý   Úcross_attention_dimÚadded_kv_proj_dimr  r°   r¸   r  rŒ   r5   r  r'  )	rý   r*  r  r°   r¸   rŒ   r5   r  r'  ©r&  r'  úgelu-approximate)r¶   Údim_outr  r   )r  r	  r   r  r   r  r   r   Únorm1_contextr>   r'   rÛ   r   r   r  Úattn2r0   ræ   r  r  rñ   Únorm2_contextÚ
ff_contextÚ_chunk_sizeÚ
_chunk_dim)
r3   r¶   r  r  r  r  r   Úcontext_norm_typer5   rD   s
            €r,   r	  zJointTransformerBlock.__init__R  s•  ø€ ô 	‰ÑÔà"4ˆÔØ 0ˆÔÙ4DÑ0È/ÐáÜ.¨sÓ3ˆD�Jä)¨#Ó.ˆDŒJàÐ 4Ò4Ü!7Ø�S¨U¸À4ÐS_ô"ˆDÕð  /Ò1Ü!1°#Ó!6ˆDÕäØ-Ð.?Ð-@Ð@~Ðóð ô ”1Ð4Ô5Ü-Ó/‰IäØkóð ô ØØ $Ø!Ø'Ø%ØØ-ØØØØô
ˆŒ	ñ Ü"ØØ$(Ø+Ø)ØØØ#ØØô
ˆD�Jð ˆDŒJä—\‘\ #¸%ÀTÔJˆŒ
Ü #¨sÐBTÔUˆŒáÜ!#§¡¨cÀeÐQUÔ!VˆDÔÜ)¨c¸3ÐN`ÔaˆD�Oà!%ˆDÔØ"ˆDŒOð  ˆÔØˆ�r-   rô   c                 ó    — || _         || _        y ré   ©r3  r4  ©r3   rô   r¶   s      r,   Úset_chunk_feed_forwardz,JointTransformerBlock.set_chunk_feed_forward¤  ó   € à%ˆÔØˆ�r-   rò   rã   ÚtembÚjoint_attention_kwargsr   c                 ób  — |xs i }| j                   r| j                  ||¬«      \  }}}}}	}
}n| j                  ||¬«      \  }}}}}	| j                  r| j                  ||«      }n| j                  ||¬«      \  }}}}} | j                  d||dœ|¤Ž\  }}|j                  d«      |z  }||z   }| j                   r- | j                  dd
i|¤Ž}j                  d«      |z  }||z   }| j                  |«      }|d|d d …d f   z   z  |d d …d f   z   }| j                  �-t        | j                  || j                  | j                  «      }n| j                  |«      }|	j                  d«      |z  }||z   }| j                  rd }||fS j                  d«      |z  }||z   }| j                  |«      }|dd d …d f   z   z  d d …d f   z   }| j                  �-t        | j                  || j                  | j                  «      }n| j                  |«      }|j                  d«      |z  z   }||fS )N)Úemb)rò   rã   r   rò   rS   )r   r  r  r/  r  rÞ   r0  r  r3  rú   rñ   r4  r1  r2  )r3   rò   rã   r;  r<  Únorm_hidden_statesÚgate_msaÚ	shift_mlpÚ	scale_mlpÚgate_mlpÚnorm_hidden_states2Ú	gate_msa2rè   Ú
c_gate_msaÚc_shift_mlpÚc_scale_mlpÚ
c_gate_mlpÚattn_outputÚcontext_attn_outputÚattn_output2rù   Úcontext_ff_outputs                         r,   r  zJointTransformerBlock.forward©  sÂ  € ð "8Ò!=¸2ÐØ×"Ò"Øko×kuÑkuØ 4ð lvó lÑhÐ ¨)°YÀÐJ]Ñ_hð LPÏ:É:ÐVcÐimÈ:ÓKnÑHÐ ¨)°YÀà× Ò Ø)-×);Ñ);Ð<QÐSWÓ)XÑ&à[_×[mÑ[mØ%¨4ð \nó \ÑXÐ&¨
°KÀÈjð
 ,5¨4¯9©9ð ,
Ø,Ø"<ñ,
ð %ñ,
Ñ(ˆÐ(ð ×(Ñ(¨Ó+¨kÑ9ˆØ%¨Ñ3ˆà×"Ò"Ø%˜4Ÿ:™:ÑbÐ4GÐbÐKaÑbˆLØ$×.Ñ.¨qÓ1°LÑ@ˆLØ)¨LÑ8ˆMà!ŸZ™Z¨Ó6ÐØ/°1°yÂÀDÀÑ7IÑ3IÑJÈYÒWXÐZ^ÐW^ÑM_Ñ_ÐØ×ÑÐ'ä-¨d¯g©gÐ7IÈ4Ï?É?Ð\`×\lÑ\lÓm‰IàŸ™Ð 2Ó3ˆIØ×&Ñ& qÓ)¨IÑ5ˆ	à%¨	Ñ1ˆð × Ò Ø$(Ð!ð  % mÐ3Ð3ð #-×"6Ñ"6°qÓ"9Ð<OÑ"OÐØ$9Ð<OÑ$OÐ!à)-×);Ñ);Ð<QÓ)RÐ&Ø)CÀqÈ;ÒWXÐZ^ÐW^ÑK_ÑG_Ñ)`ÐcnÒopÐrvÐovÑcwÑ)wÐ&Ø×ÑÐ+ä$9Ø—O‘OÐ%?ÀÇÁÐRV×RbÑRbó%Ñ!ð %)§O¡OÐ4NÓ$OÐ!Ø$9¸J×<PÑ<PÐQRÓ<SÐVgÑ<gÑ$gÐ!à$ mÐ3Ð3r-   )FNF©r   ré   )rE   rP   rQ   r  rï   rì   r.   r	  r9  r/   ÚFloatTensorr2   r   rí   rð   r  r  r  s   @r,   r  r  C  sæ   ø„ ñð$ "'Ø"Ø#(ñOàðOð !ðOð  ð	Oð
 ðOð �t‘ðOð !õOñd°°t±ð À#ó ð 9=ñC4à×(Ñ(ðC4ð  %×0Ñ0ðC4ð ×Ñð	C4ð
 !% S¨# X¡°Ñ 5ðC4ð 
ˆu�|‰|˜UŸ\™\Ð)Ñ	*÷C4r-   r  c            -       óü  ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d&dededededz  dededz  d	ed
edededededededededz  dedz  dedz  dedz  dedz  dedef,ˆ fd„Zd'dedz  defd„Z		 	 	 	 	 	 	 d(de
j                  de
j                  dz  de
j                  dz  de
j                  dz  d e
j                  dz  d!eeef   d"e
j                  dz  d#eee
j                  f   dz  d$e
j                  fd%„Zˆ xZS ))ÚBasicTransformerBlockaä  
    A basic Transformer block.

    Parameters:
        dim (`int`): The number of channels in the input and output.
        num_attention_heads (`int`): The number of heads to use for multi-head attention.
        attention_head_dim (`int`): The number of channels in each head.
        dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
        cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
        activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
        num_embeds_ada_norm (:
            obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`.
        attention_bias (:
            obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
        only_cross_attention (`bool`, *optional*):
            Whether to use only cross-attention layers. In this case two cross attention layers are used.
        double_self_attention (`bool`, *optional*):
            Whether to use two self-attention layers. In this case no cross attention layers are used.
        upcast_attention (`bool`, *optional*):
            Whether to upcast the attention computation to float32. This is useful for mixed precision training.
        norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
            Whether to use learnable elementwise affine parameters for normalization.
        norm_type (`str`, *optional*, defaults to `"layer_norm"`):
            The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
        final_dropout (`bool` *optional*, defaults to False):
            Whether to apply a final dropout after the last feed-forward layer.
        attention_type (`str`, *optional*, defaults to `"default"`):
            The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`.
        positional_embeddings (`str`, *optional*, defaults to `None`):
            The type of positional embeddings to apply to.
        num_positional_embeddings (`int`, *optional*, defaults to `None`):
            The maximum number of positional embeddings to apply.
    Nr¶   r  r  r*  r  Únum_embeds_ada_normÚattention_biasÚonly_cross_attentionÚdouble_self_attentionrÇ   Únorm_elementwise_affiner(  Únorm_epsÚfinal_dropoutÚattention_typeÚpositional_embeddingsÚnum_positional_embeddingsÚ-ada_norm_continous_conditioning_embedding_dimÚada_norm_biasÚff_inner_dimÚff_biasÚattention_out_biasc           
      ó´  •— t         ‰| �  «        || _        || _        || _        || _        || _        || _        || _        |
| _	        || _
        || _        || _        |	| _        |d uxr |dk(  | _        |d uxr |dk(  | _        |dk(  | _        |dk(  | _        |dk(  | _        |dv r|€t'        d|› d|› d	�«      ‚|| _        || _        |r|€t'        d
«      ‚|dk(  rt-        ||¬«      | _        nd | _        |dk(  rt1        ||«      | _        nO|dk(  rt5        ||«      | _        n8|dk(  rt7        |||||d«      | _        nt9        j:                  |||¬«      | _        t=        ||||||	r|nd ||¬«      | _        |€|
rk|dk(  rt1        ||«      | _         n7|dk(  rt7        |||||d«      | _         nt9        j:                  |||«      | _         t=        ||
s|nd ||||||¬«      | _!        n0|dk(  rt9        j:                  |||«      | _         nd | _         d | _!        |dk(  rt7        |||||d«      | _"        n-|dv rt9        j:                  |||«      | _"        n|dk(  rd | _"        tG        ||||||¬«      | _$        |dk(  s|dk(  rtK        ||||«      | _&        |dk(  r4t9        jN                  tQ        jR                  d|«      |dz  z  «      | _*        d | _+        d| _,        y )Nr#  Úada_normÚada_norm_singler%  Úada_norm_continuous©rb  r#  ú`norm_type` is set to úw, but `num_embeds_ada_norm` is not defined. Please make sure to define `num_embeds_ada_norm` if setting `norm_type` to r&   ú\If `positional_embedding` type is defined, `num_positition_embeddings` must also be defined.Ú
sinusoidal©Úmax_seq_lengthÚrms_normr,  ©rý   r°   r  ÚdropoutrŒ   r*  rÇ   Úout_bias©rý   r*  r°   r  rn  rŒ   rÇ   ro  )r#  rb  r%  Úlayer_norm_i2vgen©rn  r  rX  Ú	inner_dimrŒ   Úgatedzgated-text-imageé   g      à?r   )-r  r	  r¶   r  r  rn  r*  r  rS  rU  rV  rZ  r[  rT  Úuse_ada_layer_norm_zeroÚuse_ada_layer_normÚuse_ada_layer_norm_singleÚuse_layer_normÚuse_ada_layer_norm_continuousr>   r(  rR  r   Ú	pos_embedr   r  r   r   r0   ræ   r   Úattn1r  r0  Únorm3r  rñ   rü   Úfuserr  r/   r   Úscale_shift_tabler3  r4  )r3   r¶   r  r  rn  r*  r  rR  rS  rT  rU  rÇ   rV  r(  rW  rX  rY  rZ  r[  r\  r]  r^  r_  r`  rD   s                           €r,   r	  zBasicTransformerBlock.__init__  s|  ø€ ô4 	‰ÑÔØˆŒØ#6ˆÔ Ø"4ˆÔØˆŒØ#6ˆÔ Ø*ˆÔØ,ˆÔØ%:ˆÔ"Ø'>ˆÔ$Ø%:ˆÔ"Ø)BˆÔ&Ø$8ˆÔ!ð )<À4Ð(GÒ'iÈYÐZiÑMiˆÔ$Ø#6¸dÐ#BÒ"_È	ÐU_ÑH_ˆÔØ)2Ð6GÑ)GˆÔ&Ø'¨<Ñ7ˆÔØ-6Ð:OÑ-OˆÔ*àÐ5Ñ5Ð:MÐ:UÜØ(¨¨ð 4KØKTÈ+ÐUVðXóð ð
 #ˆŒØ#6ˆÔ á Ð&?Ð&GÜØnóð ð ! LÒ0Ü:¸3ÐOhÔiˆD�Nà!ˆDŒNð ˜
Ò"Ü% cÐ+>Ó?ˆD�JØ˜/Ò)Ü)¨#Ð/BÓCˆD�JØÐ/Ò/Ü/ØØ=Ø'ØØØóˆD�Jô Ÿ™ cÐ>UÐ[cÔdˆDŒJäØØ%Ø'ØØÙ7KÑ 3ÐQUØ-Ø'ô	
ˆŒ
ð Ð*Ñ.Cð ˜JÒ&Ü)¨#Ð/BÓC�•
ØÐ3Ò3Ü3ØØAØ+ØØ!Øó�•
ô  Ÿ\™\¨#¨xÐ9PÓQ�”
ä"ØÙ?TÑ$7ÐZ^Ø)Ø+ØØ#Ø!1Ø+ô	ˆD�Jð Ð-Ò-ÜŸ\™\¨#¨xÐ9PÓQ�•
à!�”
ØˆDŒJð Ð-Ò-Ü/ØØ=Ø'ØØØóˆD�Jð ÐEÑEÜŸ™ c¨8Ð5LÓMˆD�JØÐ-Ò-ØˆDŒJäØØØ'Ø'Ø"Øô
ˆŒð ˜WÒ$¨Ð:LÒ(LÜ0°Ð6IÐK^Ð`rÓsˆDŒJð Ð)Ò)Ü%'§\¡\´%·+±+¸aÀÓ2EÈÈSÉÑ2PÓ%QˆDÔ"ð  ˆÔØˆ�r-   rô   c                 ó    — || _         || _        y ré   r7  r8  s      r,   r9  z,BasicTransformerBlock.set_chunk_feed_forward»  r:  r-   rò   r¿   rã   Úencoder_attention_maskÚtimestepÚcross_attention_kwargsÚclass_labelsÚadded_cond_kwargsr   c	                 ót  — |�'|j                  dd «      �t        j                  d«       |j                  d   }	| j                  dk(  r| j                  ||«      }
nì| j                  dk(  r&| j                  ||||j                  ¬«      \  }
}}}}n·| j                  dv r| j                  |«      }
n—| j                  dk(  r| j                  ||d	   «      }
nr| j                  d
k(  rX| j                  d    |j                  |	dd«      z   j                  dd¬«      \  }}}}}}| j                  |«      }
|
d|z   z  |z   }
nt        d«      ‚| j                  �| j                  |
«      }
|�|j                  «       ni }|j                  dd «      } | j                  |
f| j                  r|nd |dœ|¤Ž}| j                  dk(  rj!                  d«      |z  }n| j                  d
k(  r|z  }||z   }|j"                  dk(  r|j%                  d«      }|�| j'                  ||d   «      }| j(                  �Ë| j                  dk(  r| j+                  ||«      }
nb| j                  dv r| j+                  |«      }
nB| j                  d
k(  r|}
n0| j                  dk(  r| j+                  ||d	   «      }
nt        d«      ‚| j                  � | j                  d
k7  r| j                  |
«      }
 | j(                  |
f||dœ|¤Ž}||z   }| j                  dk(  r| j-                  ||d	   «      }
n | j                  d
k(  s| j-                  |«      }
| j                  dk(  r|
dd d …d f   z   z  d d …d f   z   }
| j                  d
k(  r| j+                  |«      }
|
dz   z  z   }
| j.                  �-t1        | j2                  |
| j4                  | j.                  «      }n| j3                  |
«      }| j                  dk(  rj!                  d«      |z  }n| j                  d
k(  r|z  }||z   }|j"                  dk(  r|j%                  d«      }|S )NrÌ   úSPassing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.r   rb  r#  )Úhidden_dtype)r%  rq  rd  Úpooled_text_embrc  ru  rÂ   r   rÆ   zIncorrect norm usedÚgligen©rã   r¿   rØ   r  )r#  r%  rq  zIncorrect norm)ÚgetrV   Úwarningr™   r(  r  r{   r  r±   rö   r>   r{  Úcopyr:   r|  rT  rÞ   rº   Úsqueezer~  r0  r  r}  r3  rú   rñ   r4  )r3   rò   r¿   rã   r�  r‚  rƒ  r„  r…  r´   r?  r@  rA  rB  rC  Ú	shift_msaÚ	scale_msaÚgligen_kwargsrJ  rù   s                       r,   r  zBasicTransformerBlock.forwardÀ  sƒ  € ð "Ð-Ø%×)Ñ)¨'°4Ó8ÐDÜ—‘ÐtÔuð #×(Ñ(¨Ñ+ˆ
à�>‰>˜ZÒ'Ø!%§¡¨M¸8Ó!DÑØ�^‰^˜Ò.ØKOÏ:É:Ø˜x¨ÀM×DWÑDWð LVó LÑHÐ ¨)°YÁð �^‰^ÐBÑBØ!%§¡¨MÓ!:ÑØ�^‰^Ð4Ò4Ø!%§¡¨MÐ;LÐM^Ñ;_Ó!`ÑØ�^‰^Ð0Ò0à×&Ñ& tÑ,¨x×/?Ñ/?À
ÈAÈrÓ/RÑRß‰e�A˜1ˆe‹oñ KˆI�y (¨I°yÀ(ð "&§¡¨MÓ!:ÐØ!3°q¸9±}Ñ!EÈ	Ñ!QÑäÐ2Ó3Ð3à�>‰>Ð%Ø!%§¡Ð0BÓ!CÐð CYÐBdÐ!7×!<Ñ!<Ô!>ÐjlÐØ.×2Ñ2°8¸TÓBˆà �d—j‘jØð
à;?×;TÒ;TÑ"7ÐZ^Ø)ñ
ð %ñ	
ˆð �>‰>˜_Ò,Ø"×,Ñ,¨QÓ/°+Ñ=‰KØ�^‰^Ð0Ò0Ø" [Ñ0ˆKà# mÑ3ˆØ×Ñ Ò"Ø)×1Ñ1°!Ó4ˆMð Ð$Ø ŸJ™J }°mÀFÑ6KÓLˆMð �:‰:Ð!Ø�~‰~ Ò+Ø%)§Z¡Z°¸xÓ%HÑ"Ø—‘Ð#WÑWØ%)§Z¡Z°Ó%>Ñ"Ø—‘Ð#4Ò4ð &3Ñ"Ø—‘Ð#8Ò8Ø%)§Z¡Z°Ð?PÐQbÑ?cÓ%dÑ"ä Ð!1Ó2Ð2à�~‰~Ð)¨d¯n©nÐ@QÒ.QØ%)§^¡^Ð4FÓ%GÐ"à$˜$Ÿ*™*Ø"ðà&;Ø5ñð )ñ	ˆKð (¨-Ñ7ˆMð �>‰>Ð2Ò2Ø!%§¡¨MÐ;LÐM^Ñ;_Ó!`ÑØ—‘Ð#4Ò4Ø!%§¡¨MÓ!:Ðà�>‰>˜_Ò,Ø!3°q¸9ÂQÈÀWÑ;MÑ7MÑ!NÐQZÒ[\Ð^bÐ[bÑQcÑ!cÐà�>‰>Ð.Ò.Ø!%§¡¨MÓ!:ÐØ!3°q¸9±}Ñ!EÈ	Ñ!QÐà×ÑÐ'ä-¨d¯g©gÐ7IÈ4Ï?É?Ð\`×\lÑ\lÓm‰IàŸ™Ð 2Ó3ˆIà�>‰>˜_Ò,Ø ×*Ñ*¨1Ó-°	Ñ9‰IØ�^‰^Ð0Ò0Ø  9Ñ,ˆIà! MÑ1ˆØ×Ñ Ò"Ø)×1Ñ1°!Ó4ˆMàÐr-   )r×   Nr  NFFFFTr%  çñhãˆµøä>FÚdefaultNNNNNTTrN  )NNNNNNN)rE   rP   rQ   r  rï   r.   rì   rÈ   r	  r9  r/   rð   Ú
LongTensorr2   r   r  r  r  s   @r,   rQ  rQ  ï  sC  ø„ ñ ðN Ø*.Ø$Ø*.Ø$Ø%*Ø&+Ø!&Ø(,Ø%ØØ#Ø'Ø,0Ø04ØDHØ$(Ø#'ØØ#'ñ1fàðfð !ðfð  ð	fð ! 4™Zðfð ðfð ! 4™Zðfð ðfð #ðfð  $ðfð ðfð "&ðfð ðfð ðfð  ð!fð" ð#fð$  # T™zð%fð& $'¨¡:ð'fð( 8;¸T±zð)fð* ˜T‘zð+fð, ˜D‘jð-fð. ð/fð0 !õ1fñP°°t±ð À#ó ð /3Ø59Ø6:Ø,0Ø15Ø04Ø<@ñxà—|‘|ðxð Ÿ™ tÑ+ðxð  %Ÿ|™|¨dÑ2ð	xð
 !&§¡¨tÑ 3ðxð ×"Ñ" TÑ)ðxð !% S¨# X¡ðxð ×&Ñ&¨Ñ-ðxð    U§\¡\Ð 1Ñ2°TÑ9ðxð 
�‰÷xr-   rQ  c            
       óL   ‡ — e Zd ZdZ	 	 d	dedededz  dedz  fˆ fd„Zd„ Zˆ xZS )
ÚLuminaFeedForwarda'  
    A feed-forward layer.

    Parameters:
        hidden_size (`int`):
            The dimensionality of the hidden layers in the model. This parameter determines the width of the model's
            hidden representations.
        intermediate_size (`int`): The intermediate dimension of the feedforward layer.
        multiple_of (`int`, *optional*): Value to ensure hidden dimension is a multiple
            of this value.
        ffn_dim_multiplier (float, *optional*): Custom multiplier for hidden
            dimension. Defaults to None.
    Nr¶   rs  Úmultiple_ofÚffn_dim_multiplierc                 ó*  •— t         ‰| �  «        |�t        ||z  «      }|||z   dz
  |z  z  }t        j                  ||d¬«      | _        t        j                  ||d¬«      | _        t        j                  ||d¬«      | _        t        «       | _	        y )Nr   F©rŒ   )
r  r	  rï   r0   rš   Úlinear_1Úlinear_2Úlinear_3r   Úsilu)r3   r¶   rs  r˜  r™  rD   s        €r,   r	  zLuminaFeedForward.__init__J  s™   ø€ ô 	‰ÑÔàÐ)ÜÐ.°Ñ:Ó;ˆIØ I°Ñ$;¸aÑ$?ÀKÑ#OÑPˆ	äŸ	™	ØØØô
ˆŒô
 Ÿ	™	ØØØô
ˆŒô
 Ÿ	™	ØØØô
ˆŒô
 “Jˆ�	r-   c                 ó„   — | j                  | j                  | j                  |«      «      | j                  |«      z  «      S ré   )r�  rŸ  rœ  rž  )r3   re   s     r,   r  zLuminaFeedForward.forwardh  s1   € Ø�}‰}˜TŸY™Y t§}¡}°QÓ'7Ó8¸4¿=¹=ÈÓ;KÑKÓLÐLr-   )é   N)	rE   rP   rQ   r  rï   rÈ   r	  r  r  r  s   @r,   r—  r—  ;  sI   ø„ ñð$ #&Ø+/ñàðð ðð ˜4‘Zð	ð
 " D™Lõö<Mr-   r—  c                   ó²   ‡ — e Zd ZdZ	 ddedededededz  f
ˆ fd„Zd	edz  fd
„Z	 ddej                  dedej                  dz  dej                  fd„Z	ˆ xZ
S )ÚTemporalBasicTransformerBlocka÷  
    A basic Transformer block for video like data.

    Parameters:
        dim (`int`): The number of channels in the input and output.
        time_mix_inner_dim (`int`): The number of channels for temporal attention.
        num_attention_heads (`int`): The number of heads to use for multi-head attention.
        attention_head_dim (`int`): The number of channels in each head.
        cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
    Nr¶   Útime_mix_inner_dimr  r  r*  c                 óÞ  •— t         ‰| �  «        ||k(  | _        t        j                  |«      | _        t        ||d¬«      | _        t        j                  |«      | _        t        |||d ¬«      | _
        |�/t        j                  |«      | _        t        ||||¬«      | _        nd | _        d | _        t        j                  |«      | _        t        |d¬«      | _        d | _        d | _        y )Nr  )r.  r  )rý   r°   r  r*  )rý   r*  r°   r  r  )r  r	  Úis_resr0   ræ   Únorm_inr  Úff_inr  r   r|  r  r0  r}  rñ   r3  r4  )r3   r¶   r¤  r  r  r*  rD   s         €r,   r	  z&TemporalBasicTransformerBlock.__init__y  så   ø€ ô 	‰ÑÔØÐ/Ñ/ˆŒä—|‘| CÓ(ˆŒô !ØØ&Ø!ô
ˆŒ
ô —\‘\Ð"4Ó5ˆŒ
ÜØ(Ø%Ø'Ø $ô	
ˆŒ
ð Ð*ô Ÿ™Ð&8Ó9ˆDŒJÜ"Ø,Ø$7Ø)Ø+ô	ˆD�Jð ˆDŒJØˆDŒJô —\‘\Ð"4Ó5ˆŒ
ÜÐ0ÀÔHˆŒð  ˆÔØˆ�r-   rô   c                 ó    — || _         d| _        y )Nr   r7  )r3   rô   Úkwargss      r,   r9  z4TemporalBasicTransformerBlock.set_chunk_feed_forward®  s   € à%ˆÔàˆ�r-   rò   Ú
num_framesrã   r   c                 óØ  — |j                   d   }|j                   \  }}}||z  }|d d d …f   j                  ||||«      }|j                  dddd«      }|j                  ||z  ||«      }|}| j                  |«      }| j                  �-t        | j                  || j                  | j                  «      }n| j                  |«      }| j                  r||z   }| j                  |«      }	| j                  |	d ¬«      }
|
|z   }| j                  �)| j                  |«      }	| j                  |	|¬«      }
|
|z   }| j                  |«      }	| j                  �-t        | j                  |	| j                  | j                  «      }n| j                  |	«      }| j                  r||z   }n|}|d d d …f   j                  ||||«      }|j                  dddd«      }|j                  ||z  ||«      }|S )Nr   r   r   r¯   )rã   )r™   r±   r²   r§  r3  rú   r¨  r4  r¦  r  r|  r0  r  r}  rñ   )r3   rò   r«  rã   r´   Úbatch_framesÚ
seq_lengthÚchannelsÚresidualr?  rJ  rù   s               r,   r  z%TemporalBasicTransformerBlock.forward´  sú  € ð #×(Ñ(¨Ñ+ˆ
à-:×-@Ñ-@Ñ*ˆ�j (Ø! ZÑ/ˆ
à% dªA gÑ.×6Ñ6°zÀ:ÈzÐ[cÓdˆØ%×-Ñ-¨a°°A°qÓ9ˆØ%×-Ñ-¨j¸:Ñ.EÀzÐS[Ó\ˆà ˆØŸ™ ]Ó3ˆà×ÑÐ'Ü1°$·*±*¸mÈTÏ_É_Ð^b×^nÑ^nÓo‰Mà ŸJ™J }Ó5ˆMà�;Š;Ø)¨HÑ4ˆMà!ŸZ™Z¨Ó6ÐØ—j‘jÐ!3È4�jÓPˆØ# mÑ3ˆð �:‰:Ð!Ø!%§¡¨MÓ!:ÐØŸ*™*Ð%7ÐOd˜*ÓeˆKØ'¨-Ñ7ˆMð "ŸZ™Z¨Ó6Ðà×ÑÐ'Ü-¨d¯g©gÐ7IÈ4Ï?É?Ð\`×\lÑ\lÓm‰IàŸ™Ð 2Ó3ˆIà�;Š;Ø%¨Ñ5‰Mà%ˆMà% dªA gÑ.×6Ñ6°zÀ:ÈzÐ[cÓdˆØ%×-Ñ-¨a°°A°qÓ9ˆØ%×-Ñ-¨j¸:Ñ.EÀzÐS[Ó\ˆàÐr-   ré   )rE   rP   rQ   r  rï   r	  r9  r/   rð   r  r  r  s   @r,   r£  r£  l  s˜   ø„ ñ	ð" +/ñ3àð3ð  ð3ð !ð	3ð
  ð3ð ! 4™Zõ3ðj°°t±ó ð 6:ñ	7à—|‘|ð7ð ð7ð  %Ÿ|™|¨dÑ2ð	7ð
 
�‰÷7r-   r£  c                   óV   ‡ — e Zd Z	 	 	 	 ddededededededz  ded	efˆ fd
„Zd„ Zˆ xZS )ÚSkipFFTransformerBlockNr¶   r  r  Úkv_input_dimÚkv_input_dim_proj_use_biasr*  rS  r`  c
           	      ó  •— t         ‰
| �  «        ||k7  rt        j                  |||«      | _        nd | _        t        |d«      | _        t        |||||||	¬«      | _        t        |d«      | _	        t        |||||||	¬«      | _
        y )Nr$  )rý   r°   r  rn  rŒ   r*  ro  )rý   r*  r°   r  rn  rŒ   ro  )r  r	  r0   rš   Ú	kv_mapperr   r  r   r|  r  r0  )r3   r¶   r  r  r³  r´  rn  r*  rS  r`  rD   s             €r,   r	  zSkipFFTransformerBlock.__init__ï  s”   ø€ ô 	‰ÑÔØ˜3ÒÜŸY™Y |°SÐ:TÓUˆD�Nà!ˆDŒNä˜S %Ó(ˆŒ
äØØ%Ø'ØØØ 3Ø'ô
ˆŒ
ô ˜S %Ó(ˆŒ
äØØ 3Ø%Ø'ØØØ'ô
ˆ�
r-   c                 ó:  — |�|j                  «       ni }| j                  �$| j                  t        j                  |«      «      }| j	                  |«      } | j
                  |fd|i|¤Ž}||z   }| j                  |«      } | j                  |fd|i|¤Ž}||z   }|S )Nrã   )rŽ  r¶  rÛ   rŸ  r  r|  r  r0  )r3   rò   rã   rƒ  r?  rJ  s         r,   r  zSkipFFTransformerBlock.forward  sÃ   € ØBXÐBdÐ!7×!<Ñ!<Ô!>ÐjlÐà�>‰>Ð%Ø$(§N¡N´1·6±6Ð:OÓ3PÓ$QÐ!à!ŸZ™Z¨Ó6Ðà �d—j‘jØñ
à"7ð
ð %ñ
ˆð $ mÑ3ˆà!ŸZ™Z¨Ó6Ðà �d—j‘jØñ
à"7ð
ð %ñ
ˆð $ mÑ3ˆàÐr-   )r×   NFT)rE   rP   rQ   rï   rì   r	  r  r  r  s   @r,   r²  r²  î  sn   ø„ ð Ø*.Ø$Ø#'ñ(
àð(
ð !ð(
ð  ð	(
ð
 ð(
ð %)ð(
ð ! 4™Zð(
ð ð(
ð !õ(
öTr-   r²  c            /       óæ  ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d(dedededededz  ded	edz  d
edededededededededz  dedz  dedz  dededededef.ˆ fd„Zdede	e
eef      fd„Zd)dedede	e   fd„Z	 d)dedededdfd„Zd*d edz  deddfd!„Z	 	 	 	 d+d"ej                   d#ej                   dz  d$ej                   dz  d%ej                   dz  d&eeef   dej                   fd'„Zˆ xZS ),ÚFreeNoiseTransformerBlockaœ  
    A FreeNoise Transformer block.

    Parameters:
        dim (`int`):
            The number of channels in the input and output.
        num_attention_heads (`int`):
            The number of heads to use for multi-head attention.
        attention_head_dim (`int`):
            The number of channels in each head.
        dropout (`float`, *optional*, defaults to 0.0):
            The dropout probability to use.
        cross_attention_dim (`int`, *optional*):
            The size of the encoder_hidden_states vector for cross attention.
        activation_fn (`str`, *optional*, defaults to `"geglu"`):
            Activation function to be used in feed-forward.
        num_embeds_ada_norm (`int`, *optional*):
            The number of diffusion steps used during training. See `Transformer2DModel`.
        attention_bias (`bool`, defaults to `False`):
            Configure if the attentions should contain a bias parameter.
        only_cross_attention (`bool`, defaults to `False`):
            Whether to use only cross-attention layers. In this case two cross attention layers are used.
        double_self_attention (`bool`, defaults to `False`):
            Whether to use two self-attention layers. In this case no cross attention layers are used.
        upcast_attention (`bool`, defaults to `False`):
            Whether to upcast the attention computation to float32. This is useful for mixed precision training.
        norm_elementwise_affine (`bool`, defaults to `True`):
            Whether to use learnable elementwise affine parameters for normalization.
        norm_type (`str`, defaults to `"layer_norm"`):
            The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
        final_dropout (`bool` defaults to `False`):
            Whether to apply a final dropout after the last feed-forward layer.
        attention_type (`str`, defaults to `"default"`):
            The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`.
        positional_embeddings (`str`, *optional*):
            The type of positional embeddings to apply to.
        num_positional_embeddings (`int`, *optional*, defaults to `None`):
            The maximum number of positional embeddings to apply.
        ff_inner_dim (`int`, *optional*):
            Hidden dimension of feed-forward MLP.
        ff_bias (`bool`, defaults to `True`):
            Whether or not to use bias in feed-forward MLP.
        attention_out_bias (`bool`, defaults to `True`):
            Whether or not to use bias in attention output project layer.
        context_length (`int`, defaults to `16`):
            The maximum number of frames that the FreeNoise block processes at once.
        context_stride (`int`, defaults to `4`):
            The number of frames to be skipped before starting to process a new batch of `context_length` frames.
        weighting_scheme (`str`, defaults to `"pyramid"`):
            The weighting scheme to use for weighting averaging of processed latent frames. As described in the
            Equation 9. of the [FreeNoise](https://huggingface.co/papers/2310.15169) paper, "pyramid" is the default
            setting used.
    Nr¶   r  r  rn  r*  r  rR  rS  rT  rU  rÇ   rV  r(  rW  rX  rZ  r[  r^  r_  r`  Úcontext_lengthÚcontext_strideÚweighting_schemec           
      ó~  •— t         ‰| �  «        || _        || _        || _        || _        || _        || _        || _        |
| _	        || _
        || _        || _        |	| _        | j                  |||«       |d uxr |dk(  | _        |d uxr |dk(  | _        |dk(  | _        |dk(  | _        |dk(  | _        |dv r|€t)        d|› d|› d	�«      ‚|| _        || _        |r|€t)        d
«      ‚|dk(  rt/        ||¬«      | _        nd | _        t3        j4                  |||¬«      | _        t9        ||||||	r|nd ||¬«      | _        |€|
r8t3        j4                  |||«      | _        t9        ||
s|nd ||||||¬«      | _        tA        ||||||¬«      | _!        t3        j4                  |||«      | _"        d | _#        d| _$        y )Nr#  rb  rc  r%  rd  re  rf  rg  r&   rh  ri  rj  r,  rm  rp  rr  r   )%r  r	  r¶   r  r  rn  r*  r  rS  rU  rV  rZ  r[  rT  Úset_free_noise_propertiesrv  rw  rx  ry  rz  r>   r(  rR  r   r{  r0   ræ   r  r   r|  r  r0  r  rñ   r}  r3  r4  )r3   r¶   r  r  rn  r*  r  rR  rS  rT  rU  rÇ   rV  r(  rW  rX  rZ  r[  r^  r_  r`  rº  r»  r¼  rD   s                           €r,   r	  z"FreeNoiseTransformerBlock.__init__n  s1  ø€ ô4 	‰ÑÔØˆŒØ#6ˆÔ Ø"4ˆÔØˆŒØ#6ˆÔ Ø*ˆÔØ,ˆÔØ%:ˆÔ"Ø'>ˆÔ$Ø%:ˆÔ"Ø)BˆÔ&Ø$8ˆÔ!à×&Ñ& ~°~ÐGWÔXð )<À4Ð(GÒ'iÈYÐZiÑMiˆÔ$Ø#6¸dÐ#BÒ"_È	ÐU_ÑH_ˆÔØ)2Ð6GÑ)GˆÔ&Ø'¨<Ñ7ˆÔØ-6Ð:OÑ-OˆÔ*àÐ5Ñ5Ð:MÐ:UÜØ(¨¨ð 4KØKTÈ+ÐUVðXóð ð
 #ˆŒØ#6ˆÔ á Ð&?Ð&GÜØnóð ð ! LÒ0Ü:¸3ÐOhÔiˆD�Nà!ˆDŒNô —\‘\ #Ð:QÐW_Ô`ˆŒ
äØØ%Ø'ØØÙ7KÑ 3ÐQUØ-Ø'ô	
ˆŒ
ð Ð*Ñ.CÜŸ™ c¨8Ð5LÓMˆDŒJä"ØÙ?TÑ$7ÐZ^Ø)Ø+ØØ#Ø!1Ø+ô	ˆDŒJô ØØØ'Ø'Ø"Øô
ˆŒô —\‘\ # xÐ1HÓIˆŒ
ð  ˆÔØˆ�r-   r«  r   c                 ó¾   — g }t        d|| j                  z
  dz   | j                  «      D ]0  }|}t        ||| j                  z   «      }|j	                  ||f«       Œ2 |S )Nr   r   )Úrangerº  r»  ÚminÚappend)r3   r«  Úframe_indicesÚiÚwindow_startÚ
window_ends         r,   Ú_get_frame_indicesz,FreeNoiseTransformerBlock._get_frame_indicesà  sh   € ØˆÜ�q˜* t×':Ñ':Ñ:¸QÑ>À×@SÑ@SÖTˆAØˆLÜ˜Z¨¨T×-@Ñ-@Ñ)@ÓAˆJØ× Ñ  ,°
Ð!;Õ<ð Uð Ðr-   c                 óÈ  — |dk(  rdg|z  }|S |dk(  r`|dz  dk(  r*|dz  }t        t        d|dz   «      «      }||d d d…   z   }|S |dz   dz  }t        t        d|«      «      }||gz   |d d d…   z   }|S |dk(  r^|dz  dk(  r-|dz  }d	g|dz
  z  |gz   }|t        t        |dd«      «      z   }|S |dz   dz  }d	g|z  }|t        t        |dd«      «      z   }|S t        d
|› �«      ‚)NÚflatg      ð?Úpyramidr   r   r   rÂ   Údelayed_reverse_sawtoothg{®Gáz„?z'Unsupported value for weighting_scheme=)ÚlistrÀ  r>   )r3   r«  r¼  ÚweightsÚmids        r,   Ú_get_frame_weightsz,FreeNoiseTransformerBlock._get_frame_weightsè  sP  € Ø˜vÒ%Ø�e˜jÑ(ˆGð8 ˆð5  Ò*Ø˜A‰~ Ò"à  A‘o�Üœu Q¨¨a©Ó0Ó1�Ø! G©D¨b¨D¡MÑ1�ð* ˆð% " A‘~¨!Ñ+�Üœu Q¨›}Ó-�Ø! S E™/¨G±D°b°D©MÑ9�ð  ˆð Ð!;Ò;Ø˜A‰~ Ò"à  A‘o�Ø˜& C¨!¡GÑ,°¨uÑ4�Ø!¤D¬¨s°A°rÓ):Ó$;Ñ;�ð ˆð " A‘~¨!Ñ+�Ø˜& 3™,�Ø!¤D¬¨s°A°rÓ):Ó$;Ñ;�ð ˆô ÐFÐGWÐFXÐYÓZÐZr-   c                 ó.   — || _         || _        || _        y ré   )rº  r»  r¼  )r3   rº  r»  r¼  s       r,   r¾  z3FreeNoiseTransformerBlock.set_free_noise_properties  s   € ð -ˆÔØ,ˆÔØ 0ˆÕr-   rô   c                 ó    — || _         || _        y ré   r7  r8  s      r,   r9  z0FreeNoiseTransformerBlock.set_chunk_feed_forward  r:  r-   rò   r¿   rã   r�  rƒ  c                 óN  — |�'|j                  dd «      �t        j                  d«       |�|j                  «       ni }|j                  }|j
                  }	|j                  d«      }
| j                  |
«      }| j                  | j                  | j                  «      }t        j                  |||	¬«      j                  d«      j                  d«      }|d   d   |
k(  }|sU|
| j                  k  rt        d|
›d| j                  ›�«      ‚|
|d   d   z
  }|j                  |
| j                  z
  |
f«       t        j                   d|
df|¬	«      }t        j"                  |«      }t%        |«      D �]—  \  }\  }}t        j&                  |d d …||…f   «      }||z  }|d d …||…f   }| j)                  |«      }| j*                  �| j+                  |«      } | j,                  |f| j.                  r|nd |d
œ|¤Ž}||z   }|j0                  dk(  r|j3                  d«      }| j4                  �X| j7                  |«      }| j*                  � | j8                  dk7  r| j+                  |«      } | j4                  |f||d
œ|¤Ž}||z   }|t;        |«      dz
  k(  rK|sI|d d … d …fxx   |d d …| d …f   |d d …| d …f   z  z  cc<   |d d …| d …fxx   |d d …| f   z  cc<   �Œo|d d …||…fxx   ||z  z  cc<   |d d …||…fxx   |z  cc<   �Œš t        j<                  t?        |jA                  | j                  d¬«      |jA                  | j                  d¬«      «      D ��cg c]"  \  }}t        jB                  |dkD  ||z  |«      ‘Œ$ c}}d¬«      jE                  |	«      }| jG                  |«      }| jH                  �-tK        | jL                  || jN                  | jH                  «      }n| jM                  |«      }||z   }|j0                  dk(  r|j3                  d«      }|S c c}}w )NrÌ   r‡  r   ry   r   rÂ   zExpected num_frames=z1 to be greater or equal than self.context_length=)rz   r‹  rØ   rc  rÆ   )(rŒ  rV   r�  rŽ  rz   r{   ÚsizerÇ  rÏ  rº  r¼  r/   r­   rÞ   r>   rÂ  rÚ   Ú
zeros_likeÚ	enumerateÚ	ones_liker  r{  r|  rT  rº   r�  r0  r  r(  r<   r–   ÚzipÚsplitÚwhererÏ   r}  r3  rú   rñ   r4  )r3   rò   r¿   rã   r�  rƒ  Úargsrª  rz   r{   r«  rÃ  Úframe_weightsÚis_last_frame_batch_completeÚlast_frame_batch_lengthÚnum_times_accumulatedÚaccumulated_valuesrÄ  Úframe_startÚ	frame_endrÍ  Úhidden_states_chunkr?  rJ  Úaccumulated_splitÚnum_times_splitrù   s                              r,   r  z!FreeNoiseTransformerBlock.forward  s¤  € ð "Ð-Ø%×)Ñ)¨'°4Ó8ÐDÜ—‘ÐtÔuàBXÐBdÐ!7×!<Ñ!<Ô!>ÐjlÐð ×%Ñ%ˆØ×#Ñ#ˆà"×'Ñ'¨Ó*ˆ
Ø×/Ñ/°
Ó;ˆØ×/Ñ/°×0CÑ0CÀT×EZÑEZÓ[ˆÜŸ™ ]¸6ÈÔO×YÑYÐZ[Ó\×fÑfÐgiÓjˆØ'4°RÑ'8¸Ñ';¸zÑ'IÐ$ñ
 ,Ø˜D×/Ñ/Ò/Ü Ð#8¨Z¨MÐ9kÐW[×WjÑWjÐVlÐ!mÓnÐnØ&0°=ÀÑ3DÀQÑ3GÑ&GÐ#Ø× Ñ  *¨t×/BÑ/BÑ"BÀJÐ!OÔPä %§¡¨Q°
¸AÐ,>ÀvÔ NÐÜ"×-Ñ-¨mÓ<Ðä+4°]×+CÑ'ˆAÑ'�˜Yô —o‘oÐ&;ºA¸{È9Ð?TÐ<TÑ&UÓVˆGØ�}Ñ$ˆGà"/²°;¸yÐ3HÐ0HÑ"IÐð "&§¡Ð,?Ó!@Ðà�~‰~Ð)Ø%)§^¡^Ð4FÓ%GÐ"à$˜$Ÿ*™*Ø"ðà?C×?XÒ?XÑ&;Ð^bØ-ñð )ñ	ˆKð #.Ð0CÑ"CÐØ"×'Ñ'¨1Ò,Ø&9×&AÑ&AÀ!Ó&DÐ#ð �z‰zÐ%Ø%)§Z¡ZÐ0CÓ%DÐ"à—>‘>Ð-°$·.±.ÐDUÒ2UØ)-¯©Ð8JÓ)KÐ&à(˜dŸj™jØ&ðà*?Ø#9ñð -ñ	�ð '2Ð4GÑ&GÐ#à”C˜Ó&¨Ñ*Ò*Ñ3OØ"¢1Ð'>Ð&>Ñ&?Ð#?Ó@Ø'ªÐ,CÐ+CÑ+DÐ(DÑEÈÒPQÐTkÐSkÑSlÐPlÑHmÑmñÓ@ð &¢aÐ*AÐ)AÑ)BÐ&BÓCÀwÊqÐSjÐRjÐOjÑGkÑkÕCà"¢1 k°)Ð&;Ð#;Ó<Ð@SÐV]Ñ@]Ñ]Ó<Ø%¢a¨°YÐ)>Ð&>Ó?À7ÑJÕ?ðc ,Dô| Ÿ	™	ô ;>Ø&×,Ñ,¨T×-@Ñ-@ÀaÐ,ÓHØ)×/Ñ/°×0CÑ0CÈÐ/ÓKô;ôñ;Ñ6Ð% ô —‘˜O¨aÑ/Ð1BÀ_Ñ1TÐVgÕhð;òð ô	
÷ ‰"ˆU‹)ð 	ð "ŸZ™Z¨Ó6Ðà×ÑÐ'Ü-¨d¯g©gÐ7IÈ4Ï?É?Ð\`×\lÑ\lÓm‰IàŸ™Ð 2Ó3ˆIà! MÑ1ˆØ×Ñ Ò"Ø)×1Ñ1°!Ó4ˆMàÐùó-s   Í'P!
)r×   Nr  NFFFFTr%  r“  FNNNTTé   rØ   rÊ  )rÊ  rN  )NNNN)rE   rP   rQ   r  rï   rÈ   r.   rì   r	  rÌ  rí   rÇ  rÏ  r¾  r9  r/   rð   r2   r   r  r  r  s   @r,   r¹  r¹  6  sh  ø„ ñ4ðv Ø*.Ø$Ø*.Ø$Ø%*Ø&+Ø!&Ø(,Ø%ØØ#Ø,0Ø04Ø#'ØØ#'Ø ØØ )ñ1pàðpð !ðpð  ð	pð
 ðpð ! 4™Zðpð ðpð ! 4™Zðpð ðpð #ðpð  $ðpð ðpð "&ðpð ðpð ðpð  ð!pð"  # T™zð#pð$ $'¨¡:ð%pð& ˜D‘jð'pð( ð)pð* !ð+pð, ð-pð. ð/pð0 õ1pðd¨Sð °T¸%ÀÀSÀ¹/Ñ5Jó ñ¨Sð ÀCð ÐX\Ð]bÑXcó ðB QZñ1Ø!ð1Ø36ð1ØJMð1à	ó1ñ°°t±ð À#ð Èdó ð /3Ø59Ø6:Ø15ñ{à—|‘|ð{ð Ÿ™ tÑ+ð{ð  %Ÿ|™|¨dÑ2ð	{ð
 !&§¡¨tÑ 3ð{ð !% S¨# X¡ð{ð 
�‰÷{r-   r¹  c                   óŽ   ‡ — e Zd ZdZ	 	 	 	 	 	 	 ddededz  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 )r  aª  
    A feed-forward layer.

    Parameters:
        dim (`int`): The number of channels in the input.
        dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
        mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
        dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
        activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
        final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
        bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
    Nr¶   r.  Úmultrn  r  rX  rŒ   c	                 óÈ  •— t         ‰
| �  «        |€t        ||z  «      }|�|n|}|dk(  rt        |||¬«      }	|dk(  rt        ||d|¬«      }	nP|dk(  rt	        |||¬«      }	n<|dk(  rt        |||¬«      }	n(|dk(  rt        |||¬«      }	n|d	k(  rt        |||d
¬«      }	t        j                  g «      | _
        | j                  j                  	«       | j                  j                  t        j                  |«      «       | j                  j                  t        j                  |||¬«      «       |r/| j                  j                  t        j                  |«      «       y y )NÚgelur›  r-  r  )ÚapproximaterŒ   r  zgeglu-approximateÚswigluzlinear-silurŸ  )rŒ   Ú
activation)r  r	  rï   r   r   r   r   r   r0   Ú
ModuleListÚnetrÂ  ÚDropoutrš   )r3   r¶   r.  rç  rn  r  rX  rs  rŒ   Úact_fnrD   s             €r,   r	  zFeedForward.__init__   s0  ø€ ô 	‰ÑÔØÐÜ˜C $™J›ˆIØ$Ð0‘'°cˆà˜FÒ"Ü˜#˜y¨tÔ4ˆFØÐ.Ò.Ü˜#˜y°fÀ4ÔH‰FØ˜gÒ%Ü˜3 	°Ô5‰FØÐ1Ò1Ü$ S¨)¸$Ô?‰FØ˜hÒ&Ü˜C °Ô6‰FØ˜mÒ+Ü% c¨9¸4ÈFÔSˆFä—=‘= Ó$ˆŒà�‰�‰˜Ôà�‰�‰œŸ
™
 7Ó+Ô,à�‰�‰œŸ	™	 )¨W¸4Ô@ÔAáØ�H‰H�O‰OœBŸJ™J wÓ/Õ0ð r-   rò   r   c                 ó–   — t        |«      dkD  s|j                  dd «      �d}t        dd|«       | j                  D ]
  } ||«      }Œ |S )Nr   rÌ   zð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.0)r<   rŒ  r   rî  )r3   rò   rÚ  rª  Údeprecation_messager!   s         r,   r  zFeedForward.forwardÈ  sQ   € Üˆt‹9�qŠ=˜FŸJ™J w°Ó5ÐAð #UÐÜ�g˜wÐ(;Ô<Ø—h”hˆFÙ" =Ó1‰Mð àÐr-   )NrØ   r×   r  FNT)rE   rP   rQ   r  rï   rÈ   r.   rì   r	  r/   rð   r  r  r  s   @r,   r  r  ’  sˆ   ø„ ñð  #ØØØ$Ø#ØØñ&1àð&1ð �t‘ð&1ð ð	&1ð
 ð&1ð ð&1ð ð&1ð õ&1ðP U§\¡\ð ÀuÇ|Á|÷ r-   r  )8Útypingr   r   r/   Útorch.nnr0   Útorch.nn.functionalÚ
functionalrÛ   Úutilsr   r   Úutils.import_utilsr   r	   r
   Úutils.torch_utilsr   Úactivationsr   r   r   r   r   r   Úattention_processorr   r   r   Ú
embeddingsr   Únormalizationr   r   r   r   r   rv   r€   Ú
get_loggerrE   rV   r   rG   r1   rð   rï   rú   rü   r  rQ  r—  r£  r²  r¹  r  rS   r-   r,   Ú<module>rÿ     s‚  ð÷ !ã Ý ß Ð ç &ß fÑ fÝ 4ß Y× Yß UÑ UÝ 5ß qÕ qñ ÔÜà€Dð 
ˆ×	Ñ	˜HÓ	%€÷O,ñ O,÷dN%ñ N%ðb˜bŸi™ið ¸¿¹ð ÐQTð Ðbeó ð ô&˜bŸi™ió &ó ð&ðR ôh4˜BŸI™Ió h4ó ðh4ðV ôH˜BŸI™Ió Hó ðHôV
.M˜Ÿ	™	ô .Mðb ô~ B§I¡Ió ~ó ð~ôBE˜RŸY™Yô EðP ôX §	¡	ó Xó ðXôv
<�"—)‘)õ <r-   