Ë
    (täiÆ_  ã                   óš  — d dl Z d dlZd dlmZ d dlmc mZ ddlmZm	Z	 ddl
mZ ddlmZ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 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 G d„ dej                  «      Z G d„ dej                  «      Z e	dd «      rej"                  Zn 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 	 	 	 	 	 d2d+e!d,e"dz  d-e#d.e$d/e$d0ej                  fd1„Z%y)3é    Né   )Úis_torch_npu_availableÚis_torch_versioné   )Úget_activation)ÚCombinedTimestepLabelEmbeddingsÚ)PixArtAlphaCombinedTimestepSizeEmbeddingsc                   óÌ   ‡ — e Zd ZdZ	 	 	 	 	 ddededz  dedz  dededefˆ fd	„Z	 dd
ej                  dej                  dz  dej                  dz  dej                  fd„Z
ˆ xZS )ÚAdaLayerNorma©  
    Norm layer modified to incorporate timestep embeddings.

    Parameters:
        embedding_dim (`int`): The size of each embedding vector.
        num_embeddings (`int`, *optional*): The size of the embeddings dictionary.
        output_dim (`int`, *optional*):
        norm_elementwise_affine (`bool`, defaults to `False):
        norm_eps (`bool`, defaults to `False`):
        chunk_dim (`int`, defaults to `0`):
    NÚembedding_dimÚnum_embeddingsÚ
output_dimÚnorm_elementwise_affineÚnorm_epsÚ	chunk_dimc                 ó2  •— t         ‰| �  «        || _        |xs |dz  }|�t        j                  ||«      | _        nd | _        t        j                  «       | _        t        j                  ||«      | _	        t        j                  |dz  ||«      | _        y ©Nr   )ÚsuperÚ__init__r   ÚnnÚ	EmbeddingÚembÚSiLUÚsiluÚLinearÚlinearÚ	LayerNormÚnorm)Úselfr   r   r   r   r   r   Ú	__class__s          €úm/Volumes/fast/ai/experiments/MLX_z-image/.venv/lib/python3.12/site-packages/diffusers/models/normalization.pyr   zAdaLayerNorm.__init__(   s}   ø€ ô 	‰ÑÔà"ˆŒØÒ4 =°1Ñ#4ˆ
àÐ%Ü—|‘| N°MÓBˆD�HàˆDŒHä—G‘G“IˆŒ	Ü—i‘i ¨zÓ:ˆŒÜ—L‘L ¨q¡°(Ð<SÓTˆ�	ó    ÚxÚtimestepÚtembÚreturnc                 ó\  — | j                   �| j                  |«      }| j                  | j                  |«      «      }| j                  dk(  r/|j	                  dd¬«      \  }}|d d …d d d …f   }|d d …d d d …f   }n|j	                  dd¬«      \  }}| j                  |«      d|z   z  |z   }|S )Nr   r   ©Údimr   )r   r   r   r   Úchunkr   )r   r#   r$   r%   ÚshiftÚscales         r!   ÚforwardzAdaLayerNorm.forward?   s¬   € ð �8‰8ÐØ—8‘8˜HÓ%ˆDà�{‰{˜4Ÿ9™9 T›?Ó+ˆà�>‰>˜QÒð  Ÿ:™: a¨Q˜:Ó/‰LˆE�5Øš!˜T¢1˜*Ñ%ˆEØš!˜T¢1˜*Ñ%‰EàŸ:™: a¨Q˜:Ó/‰LˆE�5à�I‰I�a‹L˜A ™IÑ&¨Ñ.ˆØˆr"   )NNFçñhãˆµøä>r   )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚintÚboolÚfloatr   ÚtorchÚTensorr-   Ú__classcell__©r    s   @r!   r   r      s®   ø„ ñ
ð &*Ø!%Ø(-ØØñUàðUð ˜d™
ðUð ˜$‘Jð	Uð
 "&ðUð ðUð õUð0 bfñØ—‘ðØ).¯©¸Ñ)<ðØKPÏ<É<ÐZ^ÑK^ðà	�‰÷r"   r   c                   óD   — e Zd Zdej                  dej                  fd„Zy)ÚFP32LayerNormÚinputsr&   c                 óF  — |j                   }t        j                  |j                  «       | j                  | j
                  �| j
                  j                  «       nd | j                  �| j                  j                  «       nd | j                  «      j                  |«      S ©N)	ÚdtypeÚFÚ
layer_normr5   Únormalized_shapeÚweightÚbiasÚepsÚto)r   r<   Úorigin_dtypes      r!   r-   zFP32LayerNorm.forwardU   su   € Ø—|‘|ˆÜ�|‰|Ø�L‰L‹NØ×!Ñ!Ø#'§;¡;Ð#:ˆD�K‰K×ÑÔÀØ!%§¡Ð!6ˆD�I‰I�O‰OÔ¸DØ�H‰Hó
÷ ‰"ˆ\Ó
ð	r"   N)r/   r0   r1   r6   r7   r-   © r"   r!   r;   r;   T   s   „ ð˜eŸl™lð ¨u¯|©|ô r"   r;   c            	       óš   ‡ — e Zd ZdZddedededdfˆ fd„Z	 ddej                  d	ej                  dz  de
ej                  d
f   fd„Zˆ xZS )ÚSD35AdaLayerNormZeroXzÕ
    Norm layer adaptive layer norm zero (AdaLN-Zero).

    Parameters:
        embedding_dim (`int`): The size of each embedding vector.
        num_embeddings (`int`): The size of the embeddings dictionary.
    r   Ú	norm_typerD   r&   Nc                 óö   •— t         ‰| �  «        t        j                  «       | _        t        j
                  |d|z  |¬«      | _        |dk(  rt        j                  |dd¬«      | _        y t        d|› d�«      ‚)	Né	   ©rD   rA   Fç�íµ ÷Æ°>©Úelementwise_affinerE   úUnsupported `norm_type` (z-) provided. Supported ones are: 'layer_norm'.©
r   r   r   r   r   r   r   r   r   Ú
ValueError©r   r   rK   rD   r    s       €r!   r   zSD35AdaLayerNormZeroX.__init__i   sg   ø€ Ü‰ÑÔä—G‘G“IˆŒ	Ü—i‘i ¨q°=Ñ/@ÀtÔLˆŒØ˜Ò$ÜŸ™ ]ÀuÐRVÔWˆD�IäÐ8¸¸ÐCpÐqÓrÐrr"   Úhidden_statesr   .c           	      ó  — | j                  | j                  |«      «      }|j                  dd¬«      \	  }}}}}}}	}
}| j                  |«      }|d|d d …d f   z   z  |d d …d f   z   }|d|
d d …d f   z   z  |	d d …d f   z   }|||||||fS )NrM   r   r(   ©r   r   r*   r   )r   rV   r   Ú	shift_msaÚ	scale_msaÚgate_msaÚ	shift_mlpÚ	scale_mlpÚgate_mlpÚ
shift_msa2Ú
scale_msa2Ú	gate_msa2Únorm_hidden_statesÚnorm_hidden_states2s                 r!   r-   zSD35AdaLayerNormZeroX.forwards   sÇ   € ð
 �k‰k˜$Ÿ)™) C›.Ó)ˆØlo×luÑluØ�1ð mvó m
Ñiˆ	�9˜h¨	°9¸hÈ
ÐT^Ð`ið "ŸY™Y }Ó5ÐØ*¨a°)ºA¸t¸GÑ2DÑ.DÑEÈ	ÒRSÐUYÐRYÑHZÑZˆØ0°A¸
Â1ÀdÀ7Ñ8KÑ4KÑLÈzÒZ[Ð]aÐZaÑObÑbÐØ˜h¨	°9¸hÐH[Ð]fÐfÐfr"   ©rA   Tr>   )r/   r0   r1   r2   r3   Ústrr4   r   r6   r7   Útupler-   r8   r9   s   @r!   rJ   rJ   `   su   ø„ ññs cð s°cð sÐPTð sÐ`dõ sð $(ñgà—|‘|ðgð �\‰\˜DÑ ðgð 
ˆu�|‰|˜SÐ Ñ	!÷	gr"   rJ   c                   óN  ‡ — e Zd ZdZddededz  fˆ 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j                  ej                  ej                  ej                  ej                  f   fd„Zˆ xZS )ÚAdaLayerNormZeroúÕ
    Norm layer adaptive layer norm zero (adaLN-Zero).

    Parameters:
        embedding_dim (`int`): The size of each embedding vector.
        num_embeddings (`int`): The size of the embeddings dictionary.
    Nr   r   c                 ó^  •— t         ‰| �  «        |�t        ||«      | _        nd | _        t	        j
                  «       | _        t	        j                  |d|z  |¬«      | _        |dk(  rt	        j                  |dd¬«      | _
        y |dk(  rt        |dd¬«      | _
        y t        d	|› d
�«      ‚)Né   rN   rA   FrO   rP   Úfp32_layer_norm)rQ   rD   rR   ú@) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'.)r   r   r   r   r   r   r   r   r   r   r   r;   rT   )r   r   r   rK   rD   r    s        €r!   r   zAdaLayerNormZero.__init__‹   s    ø€ Ü‰ÑÔØÐ%Ü6°~À}ÓUˆD�HàˆDŒHä—G‘G“IˆŒ	Ü—i‘i ¨q°=Ñ/@ÀtÔLˆŒØ˜Ò$ÜŸ™ ]ÀuÐRVÔWˆD�IØÐ+Ò+Ü% mÈÐTYÔZˆD�IäØ+¨I¨;Ð6vÐwóð r"   r#   r$   Úclass_labelsÚhidden_dtyper   r&   c                 ó  — | j                   �| j                  |||¬«      }| j                  | j                  |«      «      }|j                  dd¬«      \  }}}}	}
}| j	                  |«      d|d d …d f   z   z  |d d …d f   z   }|||	|
|fS )N)ro   rk   r   r(   )r   r   r   r*   r   )r   r#   r$   rn   ro   r   rY   rZ   r[   r\   r]   r^   s               r!   r-   zAdaLayerNormZero.forward�   s˜   € ð �8‰8ÐØ—(‘(˜8 \À�(ÓMˆCØ�k‰k˜$Ÿ)™) C›.Ó)ˆØILÏÉÐSTÐZ[ÈÓI\ÑFˆ	�9˜h¨	°9¸hØ�I‰I�a‹L˜A 	ª!¨T¨'Ñ 2Ñ2Ñ3°iÂÀ4ÀÑ6HÑHˆØ�(˜I y°(Ð:Ð:r"   )NrA   T)NNNN)r/   r0   r1   r2   r3   r   r6   r7   Ú
LongTensorr?   rf   r-   r8   r9   s   @r!   rh   rh   ‚   sº   ø„ ññ cð ¸3À¹:õ ð* )-Ø04Ø+/Ø#'ñ;à�<‰<ð;ð —,‘, Ñ%ð;ð ×&Ñ&¨Ñ-ð	;ð
 —k‘k DÑ(ð;ð �\‰\˜DÑ ð;ð 
ˆu�|‰|˜UŸ\™\¨5¯<©<¸¿¹ÀuÇ|Á|ÐSÑ	T÷;r"   rh   c                   óä   ‡ — e Zd ZdZd	defˆ fd„Z	 d
dej                  dej                  dz  deej                  ej                  ej                  ej                  ej                  f   fd„Z	ˆ xZ
S )ÚAdaLayerNormZeroSingleri   r   c                 óö   •— t         ‰| �  «        t        j                  «       | _        t        j
                  |d|z  |¬«      | _        |dk(  rt        j                  |dd¬«      | _        y t        d|› d�«      ‚)	Né   rN   rA   FrO   rP   rR   rm   rS   rU   s       €r!   r   zAdaLayerNormZeroSingle.__init__¶   sk   ø€ Ü‰ÑÔä—G‘G“IˆŒ	Ü—i‘i ¨q°=Ñ/@ÀtÔLˆŒØ˜Ò$ÜŸ™ ]ÀuÐRVÔWˆD�IäØ+¨I¨;Ð6vÐwóð r"   Nr#   r   r&   c                 óÈ   — | j                  | j                  |«      «      }|j                  dd¬«      \  }}}| j                  |«      d|d d …d f   z   z  |d d …d f   z   }||fS )Nru   r   r(   rX   )r   r#   r   rY   rZ   r[   s         r!   r-   zAdaLayerNormZeroSingle.forwardÂ   sk   € ð
 �k‰k˜$Ÿ)™) C›.Ó)ˆØ),¯©°1¸!¨Ó)<Ñ&ˆ	�9˜hØ�I‰I�a‹L˜A 	ª!¨T¨'Ñ 2Ñ2Ñ3°iÂÀ4ÀÑ6HÑHˆØ�(ˆ{Ðr"   rd   r>   ©r/   r0   r1   r2   r3   r   r6   r7   rf   r-   r8   r9   s   @r!   rs   rs   ­   sk   ø„ ññ
 cõ 
ð $(ñà�<‰<ðð �\‰\˜DÑ ðð 
ˆu�|‰|˜UŸ\™\¨5¯<©<¸¿¹ÀuÇ|Á|ÐSÑ	T÷	r"   rs   c                   óÔ   ‡ — e Zd ZdZdededefˆ fd„Z	 ddej                  dej                  dz  d	e
ej                  ej                  ej                  ej                  f   fd
„Zˆ xZS )ÚLuminaRMSNormZerozˆ
    Norm layer adaptive RMS normalization zero.

    Parameters:
        embedding_dim (`int`): The size of each embedding vector.
    r   r   r   c                 óÌ   •— t         ‰| �  «        t        j                  «       | _        t        j
                  t        |d«      d|z  d¬«      | _        t        ||¬«      | _	        y )Ni   é   TrN   ©rE   )
r   r   r   r   r   r   Úminr   ÚRMSNormr   )r   r   r   r   r    s       €r!   r   zLuminaRMSNormZero.__init__Õ   sP   ø€ Ü‰ÑÔÜ—G‘G“IˆŒ	Ü—i‘iÜ�˜tÓ$Ø�ÑØô
ˆŒô
 ˜M¨xÔ8ˆ�	r"   Nr#   r   r&   c                 óº   — | j                  | j                  |«      «      }|j                  dd¬«      \  }}}}| j                  |«      d|d d …d f   z   z  }||||fS )Nr{   r   r(   rX   )r   r#   r   rZ   r[   r]   r^   s          r!   r-   zLuminaRMSNormZero.forwardß   sd   € ð
 �k‰k˜$Ÿ)™) C›.Ó)ˆØ36·9±9¸QÀA°9Ó3FÑ0ˆ	�8˜Y¨Ø�I‰I�a‹L˜A 	ª!¨T¨'Ñ 2Ñ2Ñ3ˆà�(˜I xÐ/Ð/r"   r>   )r/   r0   r1   r2   r3   r5   r4   r   r6   r7   rf   r-   r8   r9   s   @r!   ry   ry   Í   st   ø„ ñð9 cð 9°Uð 9ÐUYõ 9ð $(ñ	0à�<‰<ð	0ð �\‰\˜DÑ ð	0ð 
ˆu�|‰|˜UŸ\™\¨5¯<©<¸¿¹ÐEÑ	F÷		0r"   ry   c                   ó  ‡ — e Zd ZdZddedefˆ fd„Z	 	 	 ddej                  de	e
ej                  f   dz  dedz  d	ej                  dz  d
eej                  ej                  ej                  ej                  ej                  f   f
d„Zˆ xZS )ÚAdaLayerNormSingleaT  
    Norm layer adaptive layer norm single (adaLN-single).

    As proposed in PixArt-Alpha (see: https://huggingface.co/papers/2310.00426; Section 2.3).

    Parameters:
        embedding_dim (`int`): The size of each embedding vector.
        use_additional_conditions (`bool`): To use additional conditions for normalization or not.
    r   Úuse_additional_conditionsc                 óÀ   •— t         ‰| �  «        t        ||dz  |¬«      | _        t	        j
                  «       | _        t	        j                  |d|z  d¬«      | _        y )Nru   )Úsize_emb_dimr‚   rk   TrN   )	r   r   r	   r   r   r   r   r   r   )r   r   r‚   r    s      €r!   r   zAdaLayerNormSingle.__init__ö   sO   ø€ Ü‰ÑÔä<Ø¨¸Ñ(:ÐVoô
ˆŒô —G‘G“IˆŒ	Ü—i‘i ¨q°=Ñ/@ÀtÔLˆ�r"   Nr$   Úadded_cond_kwargsÚ
batch_sizero   r&   c                 óˆ   — |xs d d dœ} | j                   |fi |¤||dœ¤Ž}| j                  | j                  |«      «      |fS )N)Ú
resolutionÚaspect_ratio)r†   ro   )r   r   r   )r   r$   r…   r†   ro   Úembedded_timesteps         r!   r-   zAdaLayerNormSingle.forward   sS   € ð .Ò[ÀÐVZÑ1[ÐØ$˜DŸH™H XÑuÐ1BÐuÈzÐhtÓuÐØ�{‰{˜4Ÿ9™9Ð%6Ó7Ó8Ð:KÐKÐKr"   )F)NNN)r/   r0   r1   r2   r3   r4   r   r6   r7   Údictre   r?   rf   r-   r8   r9   s   @r!   r�   r�   ë   s­   ø„ ññM cð MÀdõ Mð =AØ!%Ø+/ñ
Là—,‘,ð
Lð    U§\¡\Ð 1Ñ2°TÑ9ð
Lð ˜$‘Jð	
Lð
 —k‘k DÑ(ð
Lð 
ˆu�|‰|˜UŸ\™\¨5¯<©<¸¿¹ÀuÇ|Á|ÐSÑ	T÷
Lr"   r�   c                   ó’   ‡ — e Zd ZdZ	 ddededededz  def
ˆ fd„Zd	ej                  d
ej                  dej                  fd„Z
ˆ xZS )ÚAdaGroupNormañ  
    GroupNorm layer modified to incorporate timestep embeddings.

    Parameters:
        embedding_dim (`int`): The size of each embedding vector.
        num_embeddings (`int`): The size of the embeddings dictionary.
        num_groups (`int`): The number of groups to separate the channels into.
        act_fn (`str`, *optional*, defaults to `None`): The activation function to use.
        eps (`float`, *optional*, defaults to `1e-5`): The epsilon value to use for numerical stability.
    Nr   Úout_dimÚ
num_groupsÚact_fnrE   c                 ó®   •— t         ‰| �  «        || _        || _        |€d | _        nt        |«      | _        t        j                  ||dz  «      | _        y r   )	r   r   r�   rE   Úactr   r   r   r   )r   r   rŽ   r�   r�   rE   r    s         €r!   r   zAdaGroupNorm.__init__  sL   ø€ ô 	‰ÑÔØ$ˆŒØˆŒàˆ>ØˆD�Hä% fÓ-ˆDŒHä—i‘i ¨w¸©{Ó;ˆ�r"   r#   r   r&   c                 ó  — | j                   r| j                  |«      }| j                  |«      }|d d …d d …d d f   }|j                  dd¬«      \  }}t        j                  || j
                  | j                  ¬«      }|d|z   z  |z   }|S )Nr   r   r(   r|   )r’   r   r*   r@   Ú
group_normr�   rE   )r   r#   r   r,   r+   s        r!   r-   zAdaGroupNorm.forward'  s~   € Ø�8Š8Ø—(‘(˜3“-ˆCØ�k‰k˜#ÓˆØ’!’Q˜˜dÐ"Ñ#ˆØ—y‘y ¨�yÓ*‰ˆˆuä�L‰L˜˜DŸO™O°·±Ô:ˆØ��U‘‰O˜eÑ#ˆØˆr"   )Nr.   )r/   r0   r1   r2   r3   re   r5   r   r6   r7   r-   r8   r9   s   @r!   r�   r�     sf   ø„ ñ	ð jnñ<Ø ð<Ø+.ð<Ø<?ð<ØILÈtÉð<Øafõ<ð	˜Ÿ™ð 	¨E¯L©Lð 	¸U¿\¹\÷ 	r"   r�   c                   ó†   ‡ — e Zd ZdZ	 	 	 	 d	dedefˆ fd„Zdej                  dej                  dej                  fd„Zˆ xZ	S )
ÚAdaLayerNormContinuousa�  
    Adaptive normalization layer with a norm layer (layer_norm or rms_norm).

    Args:
        embedding_dim (`int`): Embedding dimension to use during projection.
        conditioning_embedding_dim (`int`): Dimension of the input condition.
        elementwise_affine (`bool`, defaults to `True`):
            Boolean flag to denote if affine transformation should be applied.
        eps (`float`, defaults to 1e-5): Epsilon factor.
        bias (`bias`, defaults to `True`): Boolean flag to denote if bias should be use.
        norm_type (`str`, defaults to `"layer_norm"`):
            Normalization layer to use. Values supported: "layer_norm", "rms_norm".
    r   Úconditioning_embedding_dimc                 ó  •— t         ‰| �  «        t        j                  «       | _        t        j
                  ||dz  |¬«      | _        |dk(  rt        ||||«      | _        y |dk(  rt        |||«      | _        y t        d|› �«      ‚)Nr   rN   rA   Úrms_normúunknown norm_type )r   r   r   r   r   r   r   r   r   r~   rT   )r   r   r—   rQ   rE   rD   rK   r    s          €r!   r   zAdaLayerNormContinuous.__init__B  s   ø€ ô 	‰ÑÔÜ—G‘G“IˆŒ	Ü—i‘iÐ :¸MÈAÑ<MÐTXÔYˆŒØ˜Ò$Ü! -°Ð6HÈ$ÓOˆD�IØ˜*Ò$Ü ¨sÐ4FÓGˆD�IäÐ1°)°Ð=Ó>Ð>r"   r#   Úconditioning_embeddingr&   c                 ó
  — | j                  | j                  |«      j                  |j                  «      «      }t	        j
                  |dd¬«      \  }}| j                  |«      d|z   d d …d d d …f   z  |d d …d d d …f   z   }|S )Nr   r   r(   )r   r   rF   r?   r6   r*   r   )r   r#   r›   r   r,   r+   s         r!   r-   zAdaLayerNormContinuous.forwardZ  su   € à�k‰k˜$Ÿ)™)Ð$:Ó;×>Ñ>¸q¿w¹wÓGÓHˆÜ—{‘{ 3¨¨qÔ1‰ˆˆuØ�I‰I�a‹L˜A ™I¢q¨$² zÑ2Ñ2°Uº1¸dÂA¸:Ñ5FÑFˆØˆr"   )Tr.   TrA   )
r/   r0   r1   r2   r3   r   r6   r7   r-   r8   r9   s   @r!   r–   r–   3  sV   ø„ ñð.  ØØØñ?àð?ð %(õ?ð0˜Ÿ™ð ¸u¿|¹|ð ÐPU×P\ÑP\÷ r"   r–   c                   óŽ   ‡ — e Zd Z	 	 	 	 	 d
dedededz  fˆ fd„Zdej                  dej                  dej                  fd	„Zˆ xZS )ÚLuminaLayerNormContinuousNr   r—   rŽ   c                 ó\  •— t         ‰| �  «        t        j                  «       | _        t        j
                  |||¬«      | _        |dk(  rt        ||||«      | _        n'|dk(  rt        |||¬«      | _        nt        d|› �«      ‚d | _        |�t        j
                  |||¬«      | _        y y )NrN   rA   r™   ©rE   rQ   rš   )r   r   r   r   r   r   Úlinear_1r   r   r~   rT   Úlinear_2)	r   r   r—   rQ   rE   rD   rK   rŽ   r    s	           €r!   r   z"LuminaLayerNormContinuous.__init__c  s¢   ø€ ô 	‰ÑÔô —G‘G“IˆŒ	ÜŸ	™	Ð"<¸mÐRVÔWˆŒà˜Ò$Ü! -°Ð6HÈ$ÓOˆD�IØ˜*Ò$Ü °3ÐK]Ô^ˆD�IäÐ1°)°Ð=Ó>Ð>àˆŒØÐÜŸI™I m°WÀ4ÔHˆD�Mð r"   r#   r›   r&   c                 óø   — | j                  | j                  |«      j                  |j                  «      «      }|}| j	                  |«      d|z   d d …d d d …f   z  }| j
                  �| j                  |«      }|S ©Nr   )r¡   r   rF   r?   r   r¢   )r   r#   r›   r   r,   s        r!   r-   z!LuminaLayerNormContinuous.forwardƒ  sn   € ð �m‰m˜DŸI™IÐ&<Ó=×@Ñ@ÀÇÁÓIÓJˆØˆØ�I‰I�a‹L˜A ™I¢q¨$² zÑ2Ñ2ˆà�=‰=Ð$Ø—‘˜aÓ ˆAàˆr"   )Tr.   TrA   N)	r/   r0   r1   r3   r   r6   r7   r-   r8   r9   s   @r!   rž   rž   b  sk   ø„ ð  ØØØØ"ñIàðIð %(ðIð �t‘õIð@à�<‰<ðð !&§¡ðð 
�‰÷	r"   rž   c                   óþ   ‡ — e Zd ZdZdedefˆ fd„Z	 ddej                  dej                  dej                  dz  d	eej                  ej                  ej                  ej                  ej                  f   fd
„Z	ˆ xZ
S )Ú%CogView3PlusAdaLayerNormZeroTextImageri   r   r)   c                 ó  •— t         ‰| �  «        t        j                  «       | _        t        j
                  |d|z  d¬«      | _        t        j                  |dd¬«      | _        t        j                  |dd¬«      | _	        y )Né   TrN   Fr.   rP   )
r   r   r   r   r   r   r   r   Únorm_xÚnorm_c)r   r   r)   r    s      €r!   r   z.CogView3PlusAdaLayerNormZeroTextImage.__init__œ  s[   ø€ Ü‰ÑÔä—G‘G“IˆŒ	Ü—i‘i ¨r°C©x¸dÔCˆŒÜ—l‘l 3¸5ÀdÔKˆŒÜ—l‘l 3¸5ÀdÔKˆ�r"   Nr#   Úcontextr   r&   c                 óB  — | j                  | j                  |«      «      }|j                  dd¬«      \  }}}}}}	}
}}}}}| j                  |«      }| j	                  |«      }|d|d d …d f   z   z  |d d …d f   z   }|d|d d …d f   z   z  |
d d …d f   z   }|||||	|||||f
S )Nr¨   r   r(   )r   r   r*   r©   rª   )r   r#   r«   r   rY   rZ   r[   r\   r]   r^   Úc_shift_msaÚc_scale_msaÚ
c_gate_msaÚc_shift_mlpÚc_scale_mlpÚ
c_gate_mlpÚnormed_xÚnormed_contexts                     r!   r-   z-CogView3PlusAdaLayerNormZeroTextImage.forward¤  sÚ   € ð �k‰k˜$Ÿ)™) C›.Ó)ˆð �I‰I�b˜aˆIÓ ñ	
ØØØØØØØØØØØØà—;‘;˜q“>ˆØŸ™ WÓ-ˆØ˜˜I¢a¨ gÑ.Ñ.Ñ/°)ºA¸t¸GÑ2DÑDˆØ  A¨²A°t°GÑ(<Ñ$<Ñ=ÀÊAÈtÈGÑ@TÑTˆØ�(˜I y°(¸GÀZÐQ\Ð^iÐkuÐuÐur"   r>   rw   r9   s   @r!   r¦   r¦   “  sˆ   ø„ ñðL cð L°õ Lð $(ñ	và�<‰<ðvð —‘ðvð �\‰\˜DÑ ð	vð
 
ˆu�|‰|˜UŸ\™\¨5¯<©<¸¿¹ÀuÇ|Á|ÐSÑ	T÷vr"   r¦   c                   óÆ   ‡ — e Zd Z	 	 	 ddedededededdfˆ fd„Zd	ej                  d
ej                  dej                  de	ej                  ej                  f   fd„Z
ˆ xZS )ÚCogVideoXLayerNormZeroÚconditioning_dimr   rQ   rE   rD   r&   Nc                 óÎ   •— t         ‰| �  «        t        j                  «       | _        t        j
                  |d|z  |¬«      | _        t        j                  |||¬«      | _        y )Nrk   rN   r    )	r   r   r   r   r   r   r   r   r   )r   r·   r   rQ   rE   rD   r    s         €r!   r   zCogVideoXLayerNormZero.__init__Á  sL   ø€ ô 	‰ÑÔä—G‘G“IˆŒ	Ü—i‘iÐ 0°!°mÑ2CÈ$ÔOˆŒÜ—L‘L °CÐL^Ô_ˆ�	r"   rV   Úencoder_hidden_statesr%   c                 ó^  — | j                  | j                  |«      «      j                  dd¬«      \  }}}}}}	| j                  |«      d|z   d d …d d d …f   z  |d d …d d d …f   z   }| j                  |«      d|z   d d …d d d …f   z  |d d …d d d …f   z   }|||d d …d d d …f   |	d d …d d d …f   fS )Nrk   r   r(   rX   )
r   rV   r¹   r%   r+   r,   ÚgateÚ	enc_shiftÚ	enc_scaleÚenc_gates
             r!   r-   zCogVideoXLayerNormZero.forwardÏ  sØ   € ð >B¿[¹[ÈÏÉÐSWËÓ=Y×=_Ñ=_Ð`aÐghÐ=_Ó=iÑ:ˆˆu�d˜I y°(ØŸ	™	 -Ó0°A¸±IºqÀ$Ê¸zÑ3JÑJÈUÒSTÐVZÒ\]ÐS]ÑM^Ñ^ˆØ $§	¡	Ð*?Ó @ÀAÈ	ÁMÒSTÐVZÒ\]ÐS]ÑC^Ñ ^ÐajÒklÐnrÒtuÐkuÑavÑ vÐØÐ3°Tº!¸TÂ1¸*Ñ5EÀxÒPQÐSWÒYZÐPZÑG[Ð[Ð[r"   )Tr.   T)r/   r0   r1   r3   r4   r5   r   r6   r7   rf   r-   r8   r9   s   @r!   r¶   r¶   À  sž   ø„ ð
 $(ØØñ`àð`ð ð`ð !ð	`ð
 ð`ð ð`ð 
õ`ð\Ø"Ÿ\™\ð\ØBGÇ,Á,ð\ØV[×VbÑVbð\à	ˆu�|‰|˜UŸ\™\Ð)Ñ	*÷\r"   r¶   z>=z2.1.0c                   ó8   ‡ — e Zd ZdZddededefˆ fd„Zd„ Zˆ xZS )r   a°  
        LayerNorm with the bias parameter.

        Args:
            dim (`int`): Dimensionality to use for the parameters.
            eps (`float`, defaults to 1e-5): Epsilon factor.
            elementwise_affine (`bool`, defaults to `True`):
                Boolean flag to denote if affine transformation should be applied.
            bias (`bias`, defaults to `True`): Boolean flag to denote if bias should be use.
        rE   rQ   rD   c                 óˆ  •— t         ‰| �  «        || _        t        |t        j
                  «      r|f}t        j                  |«      | _        |ret        j                  t        j                  |«      «      | _        |r.t        j                  t        j                  |«      «      | _        y d | _        y d | _        d | _        y r>   )r   r   rE   Ú
isinstanceÚnumbersÚIntegralr6   ÚSizer)   r   Ú	ParameterÚonesrC   ÚzerosrD   ©r   r)   rE   rQ   rD   r    s        €r!   r   zLayerNorm.__init__é  s…   ø€ Ü‰GÑÔàˆDŒHä˜#œw×/Ñ/Ô0Ø�f�ä—z‘z #“ˆDŒHá!Ü Ÿl™l¬5¯:©:°c«?Ó;�”Ù>BœBŸL™L¬¯©°SÓ)9Ó:�•	È�•	à"�”Ø �•	r"   c                 ó„   — t        j                  || j                  | j                  | j                  | j
                  «      S r>   )r@   rA   r)   rC   rD   rE   )r   Úinputs     r!   r-   zLayerNorm.forwardú  s)   € Ü—<‘<  t§x¡x°·±¸d¿i¹iÈÏÉÓRÐRr"   )r.   TT©	r/   r0   r1   r2   r5   r4   r   r-   r8   r9   s   @r!   r   r   Ý  s)   ø„ ñ		ñ	! Uð 	!Àtð 	!ÐZ^õ 	!ö"	Sr"   r   c                   ó8   ‡ — e Zd ZdZddededefˆ fd„Zd„ Zˆ xZS )r~   a  
    RMS Norm as introduced in https://huggingface.co/papers/1910.07467 by Zhang et al.

    Args:
        dim (`int`): Number of dimensions to use for `weights`. Only effective when `elementwise_affine` is True.
        eps (`float`): Small value to use when calculating the reciprocal of the square-root.
        elementwise_affine (`bool`, defaults to `True`):
            Boolean flag to denote if affine transformation should be applied.
        bias (`bool`, defaults to False): If also training the `bias` param.
    rE   rQ   rD   c                 óˆ  •— t         ‰| �  «        || _        || _        t	        |t
        j                  «      r|f}t        j                  |«      | _	        d | _
        d | _        |r^t        j                  t        j                  |«      «      | _
        |r.t        j                  t        j                  |«      «      | _        y y y r>   )r   r   rE   rQ   rÁ   rÂ   rÃ   r6   rÄ   r)   rC   rD   r   rÅ   rÆ   rÇ   rÈ   s        €r!   r   zRMSNorm.__init__
  s’   ø€ Ü‰ÑÔàˆŒØ"4ˆÔä�cœ7×+Ñ+Ô,Ø�&ˆCä—:‘:˜c“?ˆŒàˆŒØˆŒ	áÜŸ,™,¤u§z¡z°#£Ó7ˆDŒKÙÜŸL™L¬¯©°SÓ)9Ó:�•	ð ð r"   c                 ó¨  — t        «       r³dd l}| j                  �[| j                  j                  t        j
                  t        j                  fv r%|j                  | j                  j                  «      }|j                  || j                  | j                  ¬«      d   }| j                  �|| j                  z   }|S |j                  }|j                  t        j                  «      j                  d«      j                  dd¬«      }|t	        j                  || j                  z   «      z  }| j                  �‡| j                  j                  t        j
                  t        j                  fv r%|j                  | j                  j                  «      }|| j                  z  }| j                  �|| j                  z   }|S |j                  |«      }|S )Nr   )Úepsilonr   éÿÿÿÿT©Úkeepdim)r   Ú	torch_npurC   r?   r6   Úfloat16Úbfloat16rF   Únpu_rms_normrE   rD   Úfloat32ÚpowÚmeanÚrsqrt)r   rV   rÓ   Úinput_dtypeÚvariances        r!   r-   zRMSNorm.forward  sw  € Ü!Ô#Ûà�{‰{Ð&à—;‘;×$Ñ$¬¯©¼¿¹Ð(GÑGØ$1×$4Ñ$4°T·[±[×5FÑ5FÓ$G�MØ%×2Ñ2°=À$Ç+Á+ÐW[×W_ÑW_Ð2Ó`ÐabÑcˆMØ�y‰yÐ$Ø -°·	±	Ñ 9�ð  Ðð (×-Ñ-ˆKØ$×'Ñ'¬¯©Ó6×:Ñ:¸1Ó=×BÑBÀ2ÈtÐBÓTˆHØ)¬E¯K©K¸À4Ç8Á8Ñ8KÓ,LÑLˆMà�{‰{Ð&à—;‘;×$Ñ$¬¯©¼¿¹Ð(GÑGØ$1×$4Ñ$4°T·[±[×5FÑ5FÓ$G�MØ -°·±Ñ ;�Ø—9‘9Ð(Ø$1°D·I±IÑ$=�Mð Ðð !.× 0Ñ 0°Ó =�àÐr"   )TFrË   r9   s   @r!   r~   r~   þ  s'   ø„ ñ	ñ; ð ;¸Dð ;Ètõ ;ö&r"   r~   c                   ó0   ‡ — e Zd Zddedefˆ fd„Zd„ Zˆ xZS )ÚMochiRMSNormrE   rQ   c                 ó  •— t         ‰| �  «        || _        t        |t        j
                  «      r|f}t        j                  |«      | _        |r.t        j                  t        j                  |«      «      | _        y d | _        y r>   )r   r   rE   rÁ   rÂ   rÃ   r6   rÄ   r)   r   rÅ   rÆ   rC   )r   r)   rE   rQ   r    s       €r!   r   zMochiRMSNorm.__init__=  s]   ø€ Ü‰ÑÔàˆŒä�cœ7×+Ñ+Ô,Ø�&ˆCä—:‘:˜c“?ˆŒáÜŸ,™,¤u§z¡z°#£Ó7ˆD�KàˆD�Kr"   c                 ó>  — |j                   }|j                  t        j                  «      j	                  d«      j                  dd¬«      }|t        j                  || j                  z   «      z  }| j                  �|| j                  z  }|j                  |«      }|S )Nr   rÐ   TrÑ   )	r?   rF   r6   r×   rØ   rÙ   rÚ   rE   rC   )r   rV   rÛ   rÜ   s       r!   r-   zMochiRMSNorm.forwardL  s†   € Ø#×)Ñ)ˆØ ×#Ñ#¤E§M¡MÓ2×6Ñ6°qÓ9×>Ñ>¸rÈ4Ð>ÓPˆØ%¬¯©°H¸t¿x¹xÑ4GÓ(HÑHˆà�;‰;Ð"Ø)¨D¯K©KÑ7ˆMØ%×(Ñ(¨Ó5ˆàÐr"   )T)r/   r0   r1   r5   r4   r   r-   r8   r9   s   @r!   rÞ   rÞ   <  s   ø„ ñ ð ¸Dõ ö	r"   rÞ   c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )ÚGlobalResponseNormzÈ
    Global response normalization as introduced in ConvNeXt-v2 (https://huggingface.co/papers/2301.00808).

    Args:
        dim (`int`): Number of dimensions to use for the `gamma` and `beta`.
    c                 óâ   •— t         ‰| �  «        t        j                  t	        j
                  ddd|«      «      | _        t        j                  t	        j
                  ddd|«      «      | _        y r¤   )r   r   r   rÅ   r6   rÇ   ÚgammaÚbeta)r   r)   r    s     €r!   r   zGlobalResponseNorm.__init__a  sL   ø€ Ü‰ÑÔÜ—\‘\¤%§+¡+¨a°°A°sÓ";Ó<ˆŒ
Ü—L‘L¤§¡¨Q°°1°cÓ!:Ó;ˆ�	r"   c                 óª   — t        j                  |ddd¬«      }||j                  dd¬«      dz   z  }| j                  ||z  z  | j                  z   |z   S )Nr   )r   r   T)Úpr)   rÒ   rÐ   )r)   rÒ   rO   )r6   r   rÙ   rä   rå   )r   r#   ÚgxÚnxs       r!   r-   zGlobalResponseNorm.forwardf  sS   € Ü�Z‰Z˜˜Q F°DÔ9ˆØ�2—7‘7˜r¨4�7Ó0°4Ñ7Ñ8ˆØ�z‰z˜Q ™VÑ$ t§y¡yÑ0°1Ñ4Ð4r"   )r/   r0   r1   r2   r   r-   r8   r9   s   @r!   râ   râ   X  s   ø„ ñô<ö
5r"   râ   c                   óf   ‡ — e Zd Zddededefˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )	ÚLpNormrç   r)   rE   c                 óL   •— t         ‰| �  «        || _        || _        || _        y r>   )r   r   rç   r)   rE   )r   rç   r)   rE   r    s       €r!   r   zLpNorm.__init__m  s#   ø€ Ü‰ÑÔàˆŒØˆŒØˆ�r"   rV   r&   c                 óp   — t        j                  || j                  | j                  | j                  ¬«      S )N)rç   r)   rE   )r@   Ú	normalizerç   r)   rE   )r   rV   s     r!   r-   zLpNorm.forwardt  s#   € Ü�{‰{˜=¨D¯F©F¸¿¹ÀdÇhÁhÔOÐOr"   )r   rÐ   gê-�™—q=)
r/   r0   r1   r3   r5   r   r6   r7   r-   r8   r9   s   @r!   rë   rë   l  s;   ø„ ñ˜#ð ¨ð °uõ ðP U§\¡\ð P°e·l±l÷ Pr"   rë   rK   Únum_featuresrE   rQ   rD   r&   c                 óÊ   — | dk(  rt        ||||¬«      }|S | dk(  rt        j                  ||||¬«      }|S | dk(  rt        j                  |||¬«      }|S t	        d| ›d�«      ‚)Nr™   )rE   rQ   rD   rA   Ú
batch_norm)rE   Úaffinez
norm_type=z is not supported.)r~   r   r   ÚBatchNorm2drT   )rK   rï   rE   rQ   rD   r   s         r!   Úget_normalizationrô   x  s‡   € ð �JÒÜ�|¨ÐASÐZ^Ô_ˆð €Kð 
�lÒ	"Ü�|‰|˜L¨cÐFXÐ_cÔdˆð
 €Kð	 
�lÒ	"Ü�~‰~˜l°Ð<NÔOˆð €Kô ˜J˜I˜<Ð'9Ð:Ó;Ð;r"   )rñ   Nr.   TT)&rÂ   r6   Útorch.nnr   Útorch.nn.functionalÚ
functionalr@   Úutilsr   r   Úactivationsr   Ú
embeddingsr   r	   ÚModuler   r   r;   rJ   rh   rs   ry   r�   r�   r–   rž   r¦   r¶   r~   rÞ   râ   rë   re   r3   r5   r4   rô   rH   r"   r!   Ú<module>rü      s±  ðó  ã Ý ß Ð ç <Ý 'ß bô6�2—9‘9ô 6ôr	�B—L‘Lô 	ôg˜BŸI™Iô gôD(;�r—y‘yô (;ôV˜RŸY™Yô ô@0˜Ÿ	™	ô 0ô<L˜Ÿ™ô LôD#�2—9‘9ô #ôL,˜RŸY™Yô ,ô^. §	¡	ô .ôb*v¨B¯I©Iô *vôZ\˜RŸY™Yô \ñ0 �D˜'Ô"Ø—‘�IôS�B—I‘Iô SôB9ˆb�i‰iô 9ô|�2—9‘9ô ô85˜Ÿ™ô 5ô(	PˆR�Y‰Yô 	Pð "Ø#ØØ#ØñØðà˜‘*ðð 
ðð ð	ð
 ðð ‡Y�Yôr"   